mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge eefd997d07 into b781d157d7
This commit is contained in:
commit
02e9456bcd
10 changed files with 540 additions and 43 deletions
|
|
@ -2,6 +2,8 @@
|
|||
## Helper utilities for token counting
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import re
|
||||
import struct
|
||||
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
|
||||
from typing import Final, Literal, cast
|
||||
|
|
@ -19,6 +21,7 @@ from litellm.constants import (
|
|||
DEFAULT_IMAGE_HEIGHT,
|
||||
DEFAULT_IMAGE_TOKEN_COUNT,
|
||||
DEFAULT_IMAGE_WIDTH,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
MAX_IMAGE_URL_DOWNLOAD_SIZE_MB,
|
||||
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES,
|
||||
MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES,
|
||||
|
|
@ -866,6 +869,20 @@ def _count_anthropic_content(
|
|||
return tokens
|
||||
|
||||
|
||||
LOCALLY_COUNTABLE_BLOCK_TYPES: Final = (
|
||||
"text",
|
||||
"image_url",
|
||||
"image",
|
||||
"document",
|
||||
"file",
|
||||
"tool_use",
|
||||
"tool_result",
|
||||
"thinking",
|
||||
"redacted_thinking",
|
||||
"tool_reference",
|
||||
)
|
||||
|
||||
|
||||
def _count_content_list(
|
||||
count_function: TokenCounterFunction,
|
||||
content_list: str
|
||||
|
|
@ -938,9 +955,7 @@ def _count_content_list(
|
|||
content_type = c.get("type", type(c).__name__) if isinstance(c, dict) else type(c).__name__
|
||||
raise ValueError(
|
||||
f"Invalid content item type: {content_type}. "
|
||||
f"Expected str or dict with 'type' field "
|
||||
f"(text, image_url, image, document, file, tool_use, tool_result, thinking, redacted_thinking, "
|
||||
f"tool_reference)."
|
||||
f"Expected str or dict with 'type' field ({', '.join(LOCALLY_COUNTABLE_BLOCK_TYPES)})."
|
||||
)
|
||||
return num_tokens
|
||||
except Exception as e:
|
||||
|
|
@ -1034,8 +1049,7 @@ def _format_type(props, indent):
|
|||
return " | ".join([f'"{item}"' for item in props["enum"]])
|
||||
return "string"
|
||||
elif type == "array":
|
||||
# items is required, OpenAI throws an error if it's missing
|
||||
return f"{_format_type(props['items'], indent)}[]"
|
||||
return f"{_format_type(props.get('items', {}), indent)}[]"
|
||||
elif type == "object":
|
||||
return f"{{\n{_format_object_parameters(props, indent + 2)}\n}}"
|
||||
elif type in ["integer", "number"]:
|
||||
|
|
@ -1049,3 +1063,54 @@ def _format_type(props, indent):
|
|||
else:
|
||||
# This is a guess, as an empty string doesn't yield the expected token count
|
||||
return "any"
|
||||
|
||||
|
||||
_INLINE_DATA_BASE64_RE: Final = re.compile(r"[A-Za-z0-9+/=_-]{16,}")
|
||||
|
||||
|
||||
_OPAQUE_BLOCK_KEYS: Final = frozenset(
|
||||
{"id", "tool_use_id", "cache_control", "signature", "encrypted_content", "encrypted_index"}
|
||||
)
|
||||
|
||||
|
||||
def _countable_value(key: object, value: object, depth: int) -> object:
|
||||
if key == "data" and isinstance(value, str) and _INLINE_DATA_BASE64_RE.fullmatch(value):
|
||||
return "<binary>"
|
||||
return _without_opaque_keys(value, depth + 1)
|
||||
|
||||
|
||||
def _without_opaque_keys(value: object, depth: int = 0) -> object:
|
||||
if depth > DEFAULT_MAX_RECURSE_DEPTH:
|
||||
return "<truncated>"
|
||||
if isinstance(value, Mapping):
|
||||
return { # mutable-ok: json.dumps input
|
||||
key: _countable_value(key, item, depth) for key, item in value.items() if key not in _OPAQUE_BLOCK_KEYS
|
||||
}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_without_opaque_keys(item, depth + 1) for item in value] # mutable-ok: json.dumps input
|
||||
return value
|
||||
|
||||
|
||||
def _countable_leaf_block(block: object) -> object:
|
||||
if not isinstance(block, Mapping) or block.get("type") in LOCALLY_COUNTABLE_BLOCK_TYPES:
|
||||
return block
|
||||
return {"type": "text", "text": json.dumps(_without_opaque_keys(block), default=str)}
|
||||
|
||||
|
||||
def _countable_block(block: object) -> object:
|
||||
if isinstance(block, Mapping) and block.get("type") == "tool_result" and isinstance(block.get("content"), list):
|
||||
return {**block, "content": [_countable_leaf_block(item) for item in block["content"]]}
|
||||
return _countable_leaf_block(block)
|
||||
|
||||
|
||||
def _countable_message(message: object) -> object:
|
||||
if not isinstance(message, Mapping) or not isinstance(message.get("content"), list):
|
||||
return message
|
||||
return {
|
||||
**message,
|
||||
"content": [_countable_block(block) for block in message["content"]],
|
||||
}
|
||||
|
||||
|
||||
def messages_with_uncountable_blocks_as_text(messages: Sequence[object]) -> tuple[object, ...]:
|
||||
return tuple(_countable_message(message) for message in messages)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import datetime
|
|||
import json
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -15,6 +15,9 @@ from litellm.secret_managers.main import get_secret_str
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
GEMINI_IMAGE_ASPECT_RATIOS: Final[dict[str, float]] = {
|
||||
"1:1": 1 / 1,
|
||||
"1:4": 1 / 4,
|
||||
|
|
@ -491,11 +494,26 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
request_model: str = "",
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
system: object | None = None,
|
||||
client: "httpx.AsyncClient | AsyncHTTPHandler | None" = None,
|
||||
) -> TokenCountResponse | None:
|
||||
import copy
|
||||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
||||
def failed(
|
||||
message: str, status_code: int, original_response: dict[str, object] | None = None
|
||||
) -> TokenCountResponse:
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message=message,
|
||||
status_code=status_code,
|
||||
original_response=original_response,
|
||||
)
|
||||
|
||||
deployment = deployment or {}
|
||||
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
|
||||
count_tokens_params: Final = {
|
||||
|
|
@ -503,17 +521,24 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
"contents": contents,
|
||||
}
|
||||
count_tokens_params_request.update(count_tokens_params)
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**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=result.get("tokenizer_used", ""),
|
||||
original_response=result,
|
||||
try:
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
client=client,
|
||||
**count_tokens_params_request,
|
||||
)
|
||||
|
||||
return None
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
return failed(e.message, e.status_code)
|
||||
total_tokens: Final = result.get("totalTokens") if isinstance(result, dict) else None
|
||||
if not isinstance(total_tokens, int) or isinstance(total_tokens, bool):
|
||||
return failed(
|
||||
"Google Gen AI Studio countTokens response has no totalTokens",
|
||||
502,
|
||||
result if isinstance(result, dict) else None,
|
||||
)
|
||||
return TokenCountResponse(
|
||||
total_tokens=total_tokens,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type=result.get("tokenizer_used", ""),
|
||||
original_response=result,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -84,6 +84,7 @@ class GoogleAIStudioTokenCounter:
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
client: httpx.AsyncClient | AsyncHTTPHandler | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -96,6 +97,7 @@ class GoogleAIStudioTokenCounter:
|
|||
api_key: Optional Google API key (will fall back to environment)
|
||||
api_base: Optional API base URL (defaults to Google Gen AI Studio)
|
||||
timeout: Optional timeout for the request
|
||||
client: Optional HTTP client to send the request with
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
|
|
@ -113,10 +115,8 @@ class GoogleAIStudioTokenCounter:
|
|||
}
|
||||
|
||||
Raises:
|
||||
ValueError: If API key is missing
|
||||
litellm.APIError: If the API call fails
|
||||
litellm.APIConnectionError: If the connection fails
|
||||
Exception: For any other unexpected errors
|
||||
litellm.APIError: If the API returns an error status or a body that is not JSON
|
||||
litellm.APIConnectionError: If the request fails or times out
|
||||
"""
|
||||
|
||||
# Prepare headers
|
||||
|
|
@ -132,31 +132,31 @@ class GoogleAIStudioTokenCounter:
|
|||
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
|
||||
request_body: Final = {"contents": cleaned_contents}
|
||||
|
||||
async_httpx_client: Final = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GEMINI,
|
||||
)
|
||||
async_httpx_client: Final = client or get_async_httpx_client(llm_provider=LlmProviders.GEMINI)
|
||||
|
||||
try:
|
||||
response: Final = await async_httpx_client.post(url=url, headers=headers, json=request_body)
|
||||
|
||||
# Check for HTTP errors
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse response
|
||||
result: Final = response.json()
|
||||
return result
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_msg = f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}"
|
||||
raise litellm.APIError(
|
||||
message=error_msg,
|
||||
message=f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}",
|
||||
llm_provider="gemini",
|
||||
model=model,
|
||||
status_code=e.response.status_code,
|
||||
) from e
|
||||
except httpx.RequestError as e:
|
||||
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
|
||||
except (httpx.RequestError, litellm.Timeout) as e:
|
||||
raise litellm.APIConnectionError(
|
||||
message=f"Request to Google Gen AI Studio failed: {e}", llm_provider="gemini", model=model
|
||||
) from e
|
||||
|
||||
try:
|
||||
return 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=502,
|
||||
) from e
|
||||
|
|
|
|||
|
|
@ -13729,6 +13729,7 @@ async def run_thread(
|
|||
# dependencies=[Depends(user_api_key_auth)],
|
||||
# )
|
||||
# async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)):
|
||||
from litellm.litellm_core_utils.token_counter import messages_with_uncountable_blocks_as_text
|
||||
from litellm.llms.base_llm.base_utils import BaseTokenCounter
|
||||
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
|
|
@ -13841,7 +13842,8 @@ async def _try_provider_token_count(
|
|||
code=result.status_code or 500,
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Provider token counting failed (%s): %s. Falling back to local tokenizer.",
|
||||
"Provider token counting for model %s failed (%s): %s. Falling back to local tokenizer.",
|
||||
model_to_use,
|
||||
result.status_code,
|
||||
result.error_message,
|
||||
)
|
||||
|
|
@ -13958,7 +13960,8 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
|
|||
tokenizer_used: Final = str(_tokenizer_used["type"])
|
||||
system_message: Final = _system_message(system)
|
||||
typed_messages: Final = cast( # cast-ok: request messages are raw chat-shaped dicts that token_counter normalizes
|
||||
Sequence[AllMessageValues] | None, messages
|
||||
Sequence[AllMessageValues] | None,
|
||||
None if messages is None else messages_with_uncountable_blocks_as_text(messages),
|
||||
)
|
||||
counted_messages: Final = (
|
||||
typed_messages if typed_messages is None or system_message is None else (system_message, *typed_messages)
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_mask_sequence", # max depth set.
|
||||
"_delete_nested_value_custom", # max depth set (bounded by number of path segments).
|
||||
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
|
||||
"_without_opaque_keys", # max depth set (DEFAULT_MAX_RECURSE_DEPTH).
|
||||
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.
|
||||
"_validate_inheritance_chain", # max depth set (default 100) to prevent infinite recursion in policy inheritance validation.
|
||||
"_basic_json_schema_validate", # max depth set.
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm import token_counter as token_counter_old
|
|||
import litellm.constants
|
||||
from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_COUNTS
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.litellm_core_utils.token_counter import (
|
||||
_encoding_count,
|
||||
_get_exact_count_function,
|
||||
|
|
@ -35,6 +36,7 @@ from litellm.litellm_core_utils.token_counter import (
|
|||
_get_tiktoken_count_function,
|
||||
calculate_img_tokens,
|
||||
high_detail_image_token_upper_bound,
|
||||
messages_with_uncountable_blocks_as_text,
|
||||
offload_token_count,
|
||||
)
|
||||
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
|
||||
|
|
@ -1633,3 +1635,83 @@ def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_t
|
|||
"custom": expected["Xenova/llama-3-tokenizer"],
|
||||
"requested": sorted(served),
|
||||
}
|
||||
|
||||
|
||||
def test_token_counter_counts_array_parameter_without_items():
|
||||
messages = [{"role": "user", "content": "tag this"}]
|
||||
tags_tool = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "set_tags",
|
||||
"parameters": {"type": "object", "properties": {"tags": {"type": "array"}}, "required": ["tags"]},
|
||||
},
|
||||
}
|
||||
|
||||
assert token_counter(model="gpt-4o", messages=messages, tools=[tags_tool]) > token_counter(
|
||||
model="gpt-4o", messages=messages
|
||||
)
|
||||
|
||||
|
||||
class _Unprintable:
|
||||
def __str__(self) -> str:
|
||||
raise AssertionError("opaque block values must be dropped before they are serialized")
|
||||
|
||||
|
||||
def test_uncountable_block_drops_opaque_values_without_serializing_them():
|
||||
(message,) = messages_with_uncountable_blocks_as_text(
|
||||
[
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": _Unprintable(),
|
||||
"content": [{"type": "web_search_result", "title": "Paris", "encrypted_content": _Unprintable()}],
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert message["content"] == [
|
||||
{
|
||||
"type": "text",
|
||||
"text": '{"type": "web_search_tool_result", "content": [{"type": "web_search_result", "title": "Paris"}]}',
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_uncountable_block_nesting_past_the_depth_limit_is_truncated():
|
||||
nested: object = "leaf"
|
||||
for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 5):
|
||||
nested = {"child": nested}
|
||||
|
||||
(message,) = messages_with_uncountable_blocks_as_text(
|
||||
[{"role": "assistant", "content": [{"type": "server_tool_use", "input": nested}]}]
|
||||
)
|
||||
|
||||
text = message["content"][0]["text"]
|
||||
assert text.endswith('"<truncated>"' + "}" * (DEFAULT_MAX_RECURSE_DEPTH + 1))
|
||||
assert "leaf" not in text
|
||||
|
||||
|
||||
|
||||
def test_uncountable_block_elides_inline_base64_data_but_keeps_plain_text_data():
|
||||
(message,) = messages_with_uncountable_blocks_as_text(
|
||||
[
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "code_execution_tool_result",
|
||||
"content": {"data": "iVBORw0KGgo" * 20, "stdout": "ok", "notes": {"data": "two words"}},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert message["content"][0]["text"] == (
|
||||
'{"type": "code_execution_tool_result", '
|
||||
'"content": {"data": "<binary>", "stdout": "ok", "notes": {"data": "two words"}}}'
|
||||
)
|
||||
|
|
|
|||
0
tests/unit/llms/gemini/count_tokens/__init__.py
Normal file
0
tests/unit/llms/gemini/count_tokens/__init__.py
Normal file
65
tests/unit/llms/gemini/count_tokens/test_handler.py
Normal file
65
tests/unit/llms/gemini/count_tokens/test_handler.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_non_json_body_raises_api_error_with_502():
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, content=b"<html>proxy error page</html>")
|
||||
|
||||
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=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 502
|
||||
assert "non-JSON" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_lets_internal_errors_propagate():
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
raise RuntimeError("transport exploded")
|
||||
|
||||
with pytest.raises(RuntimeError, match="transport exploded"):
|
||||
await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
model="gemini-2.5-flash",
|
||||
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
|
||||
api_key="test-key",
|
||||
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
|
||||
)
|
||||
|
||||
|
||||
def _timing_out(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.ReadTimeout("timed out", request=request)
|
||||
|
||||
|
||||
def _litellm_handler_timing_out() -> AsyncHTTPHandler:
|
||||
handler = AsyncHTTPHandler()
|
||||
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(_timing_out))
|
||||
return handler
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"client",
|
||||
(
|
||||
pytest.param(httpx.AsyncClient(transport=httpx.MockTransport(_timing_out)), id="httpx-client"),
|
||||
pytest.param(_litellm_handler_timing_out(), id="litellm-http-handler"),
|
||||
),
|
||||
)
|
||||
async def test_acount_tokens_raises_connection_error_on_timeout(client):
|
||||
with pytest.raises(litellm.APIConnectionError):
|
||||
await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
model="gemini-2.5-flash",
|
||||
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
|
||||
api_key="test-key",
|
||||
client=client,
|
||||
)
|
||||
|
|
@ -89,6 +89,103 @@ class TestGeminiModelInfo:
|
|||
|
||||
|
||||
class TestGoogleAIStudioTokenCounter:
|
||||
async def _count(self, handler, litellm_params=None):
|
||||
import httpx
|
||||
|
||||
return await GoogleAIStudioTokenCounter().count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=None,
|
||||
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
|
||||
deployment={"litellm_params": litellm_params or {"api_key": "test-key"}},
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_provider_error_returns_error_response(self):
|
||||
import httpx
|
||||
|
||||
result = await self._count(
|
||||
lambda request: httpx.Response(
|
||||
400, json={"error": {"code": 400, "message": "bad request", "status": "INVALID_ARGUMENT"}}
|
||||
)
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 400
|
||||
assert result.total_tokens == 0
|
||||
assert "bad request" in (result.error_message or "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_without_api_key_returns_provider_error_response(self, monkeypatch):
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.delenv("GEMINI_API_KEY", raising=False)
|
||||
monkeypatch.delenv("GOOGLE_API_KEY", raising=False)
|
||||
monkeypatch.setattr(litellm, "api_key", None)
|
||||
recorded = []
|
||||
|
||||
def _handler(request):
|
||||
recorded.append(request)
|
||||
return httpx.Response(
|
||||
403, json={"error": {"code": 403, "message": "API key not valid", "status": "PERMISSION_DENIED"}}
|
||||
)
|
||||
|
||||
result = await self._count(_handler, litellm_params={"model": "gemini/gemini-2.5-flash"})
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 403
|
||||
assert len(recorded) == 1 and "x-goog-api-key" not in recorded[0].headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_connection_error_returns_error_response(self):
|
||||
import httpx
|
||||
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.ConnectError("connection refused", request=request)
|
||||
|
||||
result = await self._count(_handler)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 500
|
||||
assert "connection refused" in (result.error_message or "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"upstream_json",
|
||||
[
|
||||
{"totalTokens": "abc"},
|
||||
{"totalTokens": True},
|
||||
{"promptTokensDetails": []},
|
||||
[{"totalTokens": 5}],
|
||||
],
|
||||
)
|
||||
async def test_count_tokens_malformed_provider_response_returns_502(self, upstream_json):
|
||||
import httpx
|
||||
|
||||
result = await self._count(lambda request: httpx.Response(200, json=upstream_json))
|
||||
|
||||
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 "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_valid_response_returns_the_count(self):
|
||||
import httpx
|
||||
|
||||
result = await self._count(lambda request: httpx.Response(200, json={"totalTokens": 7}))
|
||||
|
||||
assert result is not None
|
||||
assert result.error is not True
|
||||
assert result.total_tokens == 7
|
||||
|
||||
"""Test suite for GoogleAIStudioTokenCounter class"""
|
||||
|
||||
def test_should_use_token_counting_api(self):
|
||||
|
|
@ -158,7 +255,7 @@ 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
|
||||
)
|
||||
|
||||
def test_clean_contents_for_gemini_api_removes_id_field(self):
|
||||
|
|
|
|||
|
|
@ -1245,3 +1245,162 @@ async def test_anthropic_endpoint_429_rate_limit_error_format():
|
|||
finally:
|
||||
anthropic_endpoints._read_request_body = original_read_request_body
|
||||
proxy_server.token_counter = original_token_counter
|
||||
|
||||
|
||||
def _server_tool_history(stdout: str, encrypted_content: str) -> list[dict[str, object]]:
|
||||
return [
|
||||
{"role": "user", "content": "weather in Paris?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "paris"}},
|
||||
{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [
|
||||
{
|
||||
"type": "web_search_result",
|
||||
"url": "https://example.com/paris",
|
||||
"title": "Paris weather",
|
||||
"encrypted_content": encrypted_content,
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_2",
|
||||
"content": {"type": "bash_code_execution_result", "stdout": stdout, "stderr": "", "return_code": 0},
|
||||
},
|
||||
{
|
||||
"type": "text_editor_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_3",
|
||||
"content": {"type": "text_editor_code_execution_view_result", "content": "notes"},
|
||||
},
|
||||
{"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_1",
|
||||
"content": [
|
||||
{
|
||||
"type": "search_result",
|
||||
"source": "https://example.com",
|
||||
"title": "t",
|
||||
"content": [{"type": "text", "text": "18C"}],
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_token_count_estimates_server_tool_history_without_counting_ciphertext(monkeypatch):
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", None)
|
||||
|
||||
async def count(stdout: str, encrypted_content: str) -> int:
|
||||
result = await token_counter(
|
||||
request=TokenCountRequest(model="gpt-4o", messages=_server_tool_history(stdout, encrypted_content))
|
||||
)
|
||||
return result.total_tokens
|
||||
|
||||
baseline = await count("18C", "RW5jcnlwdGVk")
|
||||
|
||||
assert baseline > 0
|
||||
assert await count("18C", "RW5jcnlwdGVk" * 2000) == baseline
|
||||
assert await count("18C and sunny for the rest of the week", "RW5jcnlwdGVk") > baseline
|
||||
|
||||
|
||||
def _gemini_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gemini-count",
|
||||
"litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "fake-gemini-key"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
_GEMINI_COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_count_error_falls_back_to_the_local_estimate(monkeypatch, respx_mock):
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", False)
|
||||
count_route = respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(
|
||||
return_value=httpx.Response(400, json={"error": {"code": 400, "message": "API key not valid"}})
|
||||
)
|
||||
|
||||
response = await token_counter(
|
||||
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hello world"}]),
|
||||
call_endpoint=True,
|
||||
)
|
||||
|
||||
assert count_route.called
|
||||
assert response.error is not True
|
||||
assert response.total_tokens > 0
|
||||
assert response.tokenizer_type != "gemini_api"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_count_error_is_returned_when_fallback_is_disabled(monkeypatch, respx_mock):
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", True)
|
||||
respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(
|
||||
return_value=httpx.Response(400, json={"error": {"code": 400, "message": "API key not valid"}})
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await token_counter(
|
||||
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hi"}]),
|
||||
call_endpoint=True,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert "API key not valid" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_non_json_success_body_surfaces_as_bad_gateway_when_fallback_disabled(monkeypatch, respx_mock):
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", True)
|
||||
respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(return_value=httpx.Response(200, content=b"<html>portal</html>"))
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await token_counter(
|
||||
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hi"}]),
|
||||
call_endpoint=True,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "502"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_estimate_counts_a_tool_with_an_array_property_without_items(monkeypatch):
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", None)
|
||||
tags_tool = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "set_tags",
|
||||
"description": "Set tags",
|
||||
"parameters": {"type": "object", "properties": {"tags": {"type": "array"}}, "required": ["tags"]},
|
||||
},
|
||||
}
|
||||
|
||||
async def count(tools: list[dict[str, object]] | None) -> int:
|
||||
result = await token_counter(
|
||||
request=TokenCountRequest(model="gpt-4o", messages=[{"role": "user", "content": "tag this"}], tools=tools)
|
||||
)
|
||||
return result.total_tokens
|
||||
|
||||
assert await count([tags_tool]) > await count(None)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue