mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(gemini): validate Anthropic token-count input with TypeAdapter and type the countTokens request body
This commit is contained in:
parent
629a068e10
commit
78a8a4afe9
7 changed files with 271 additions and 118 deletions
|
|
@ -6,6 +6,7 @@ from collections.abc import Mapping, Sequence
|
|||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
|
@ -13,6 +14,7 @@ from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
|||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.gemini import GeminiCountTokensDeploymentParams
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
|
|
@ -472,6 +474,10 @@ def get_api_key_from_env() -> str | None:
|
|||
return get_secret_str("GOOGLE_API_KEY") or get_secret_str("GEMINI_API_KEY")
|
||||
|
||||
|
||||
_COUNT_TOKENS_DEPLOYMENT_PARAMS: Final = TypeAdapter(GeminiCountTokensDeploymentParams)
|
||||
_NO_DEPLOYMENT_PARAMS: Final[GeminiCountTokensDeploymentParams] = {}
|
||||
|
||||
|
||||
class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
||||
"""Token counter implementation for Google AI Studio provider."""
|
||||
|
||||
|
|
@ -496,31 +502,45 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
tools: list[dict[str, object]] | None = None,
|
||||
system: object | None = None,
|
||||
) -> TokenCountResponse | None:
|
||||
import copy
|
||||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
from litellm.llms.gemini.count_tokens.transformation import (
|
||||
AnthropicCountTokensInput,
|
||||
InvalidAnthropicRequest,
|
||||
build_count_tokens_payload,
|
||||
)
|
||||
|
||||
if contents is None and not messages:
|
||||
return None
|
||||
|
||||
payload: Final = (
|
||||
build_count_tokens_payload(model=model_to_use, messages=messages or [], system=system, tools=tools)
|
||||
if contents is None
|
||||
else None
|
||||
)
|
||||
deployment = deployment or {}
|
||||
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
|
||||
count_tokens_params: Final = {
|
||||
anthropic_input: Final[AnthropicCountTokensInput] = {
|
||||
"model": model_to_use,
|
||||
"contents": payload.contents if payload is not None else contents,
|
||||
**({"system_instruction": payload.system_instruction} if payload and payload.system_instruction else {}),
|
||||
**({"tools": payload.tools} if payload and payload.tools else {}),
|
||||
"messages": messages or (),
|
||||
"system": system,
|
||||
"tools": tools,
|
||||
}
|
||||
count_tokens_params_request.update(count_tokens_params)
|
||||
payload: Final = build_count_tokens_payload(anthropic_input) if contents is None else None
|
||||
if isinstance(payload, InvalidAnthropicRequest):
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message=payload.message,
|
||||
status_code=400,
|
||||
)
|
||||
deployment_params: Final = (
|
||||
_COUNT_TOKENS_DEPLOYMENT_PARAMS.validate_python(deployment["litellm_params"])
|
||||
if deployment and "litellm_params" in deployment
|
||||
else _NO_DEPLOYMENT_PARAMS
|
||||
)
|
||||
try:
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**count_tokens_params_request,
|
||||
model=model_to_use,
|
||||
api_key=deployment_params.get("api_key") or deployment_params.get("gemini_api_key"),
|
||||
api_base=deployment_params.get("api_base"),
|
||||
contents=contents if payload is None else payload.contents,
|
||||
system_instruction=None if payload is None else payload.system_instruction,
|
||||
tools=None if payload is None else payload.tools,
|
||||
client=self.client,
|
||||
)
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
|
|
|
|||
|
|
@ -1,16 +1,48 @@
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
|
||||
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:
|
||||
from litellm.types.google_genai.main import GenerateContentContentListUnionDict
|
||||
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:
|
||||
|
|
@ -86,7 +118,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: AsyncHTTPHandler | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, Any]:
|
||||
|
|
@ -135,18 +167,11 @@ class GoogleAIStudioTokenCounter:
|
|||
)
|
||||
|
||||
# Prepare request body - clean up contents to remove unsupported fields
|
||||
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
|
||||
request_body: Final = (
|
||||
{"contents": cleaned_contents}
|
||||
if system_instruction is None and tools is None
|
||||
else {
|
||||
"generateContentRequest": {
|
||||
"model": f"models/{model}",
|
||||
"contents": cleaned_contents,
|
||||
**({"systemInstruction": system_instruction} if system_instruction is not None else {}),
|
||||
**({"tools": tools} if tools is not None else {}),
|
||||
}
|
||||
}
|
||||
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(
|
||||
|
|
@ -154,7 +179,11 @@ class GoogleAIStudioTokenCounter:
|
|||
)
|
||||
|
||||
try:
|
||||
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
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_msg = f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}"
|
||||
raise litellm.APIError(
|
||||
|
|
|
|||
|
|
@ -1,5 +1,9 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final
|
||||
from typing import Final, cast
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
|
|
@ -13,41 +17,57 @@ from litellm.types.llms.anthropic import AnthropicMessagesRequest
|
|||
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
|
||||
|
||||
|
||||
class AnthropicCountTokensInput(TypedDict):
|
||||
"""The Anthropic-shaped fields of a token-count request before validation."""
|
||||
|
||||
model: ReadOnly[str]
|
||||
messages: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
system: ReadOnly[object]
|
||||
tools: ReadOnly[Sequence[Mapping[str, object]] | None]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiCountTokensPayload:
|
||||
contents: list[ContentType]
|
||||
contents: Sequence[ContentType]
|
||||
system_instruction: SystemInstructions | None
|
||||
tools: list[Tools] | None
|
||||
tools: Sequence[Tools] | None
|
||||
|
||||
|
||||
def build_count_tokens_payload(
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
system: object | None,
|
||||
tools: list[dict[str, Any]] | None,
|
||||
) -> GeminiCountTokensPayload:
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InvalidAnthropicRequest:
|
||||
message: str
|
||||
|
||||
|
||||
_ANTHROPIC_REQUEST: Final = TypeAdapter(AnthropicMessagesRequest)
|
||||
|
||||
|
||||
def _validated_request(raw: AnthropicCountTokensInput) -> AnthropicMessagesRequest | InvalidAnthropicRequest:
|
||||
try:
|
||||
_ANTHROPIC_REQUEST.validate_python(raw)
|
||||
except ValidationError as e:
|
||||
return InvalidAnthropicRequest(message=str(e))
|
||||
return cast(AnthropicMessagesRequest, raw) # cast-ok: validated above; pydantic returns lazy Iterable validators
|
||||
|
||||
|
||||
def build_count_tokens_payload(raw: AnthropicCountTokensInput) -> GeminiCountTokensPayload | InvalidAnthropicRequest:
|
||||
"""Translate an Anthropic Messages token-count request into the Gemini countTokens shape."""
|
||||
anthropic_request: Final = AnthropicMessagesRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
**({"system": system} if isinstance(system, (str, list)) else {}),
|
||||
**({"tools": tools} if tools else {}),
|
||||
)
|
||||
request: Final = _validated_request(raw)
|
||||
if isinstance(request, InvalidAnthropicRequest):
|
||||
return request
|
||||
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_request, custom_llm_provider="gemini"
|
||||
request, custom_llm_provider="gemini"
|
||||
)
|
||||
system_instruction, remaining_messages = _transform_system_message(
|
||||
supports_system_message=True,
|
||||
messages=list(openai_request["messages"]),
|
||||
supports_system_message=True, messages=openai_request["messages"]
|
||||
)
|
||||
contents: Final = _gemini_convert_messages_with_history(
|
||||
messages=remaining_messages, model=model, custom_llm_provider="gemini"
|
||||
messages=remaining_messages, model=raw["model"], custom_llm_provider="gemini"
|
||||
)
|
||||
openai_tools: Final = openai_request.get("tools")
|
||||
gemini_tools: Final = (
|
||||
VertexGeminiConfig()._map_function( # pyright: ignore[reportPrivateUsage] # same tool mapper the Gemini chat path uses
|
||||
value=[dict(tool) for tool in openai_tools],
|
||||
optional_params={},
|
||||
value=[dict(tool) for tool in openai_tools], # mutable-ok: _map_function only accepts list[dict]
|
||||
optional_params={}, # mutable-ok: _map_function writes retrieval config into the dict it is given
|
||||
)
|
||||
if openai_tools
|
||||
else None
|
||||
|
|
|
|||
|
|
@ -1,17 +1,42 @@
|
|||
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 GeminiCountTokensDeploymentParams(TypedDict, total=False):
|
||||
"""The deployment litellm_params the countTokens handler reads: everything else is ignored."""
|
||||
|
||||
api_key: ReadOnly[str | None]
|
||||
gemini_api_key: ReadOnly[str | None]
|
||||
api_base: ReadOnly[str | None]
|
||||
|
||||
|
||||
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,6 @@ from litellm.proxy._types import ProxyException, TokenCountRequest
|
|||
from litellm.proxy.anthropic_endpoints.endpoints import (
|
||||
count_tokens as anthropic_count_tokens,
|
||||
)
|
||||
from litellm.llms.base_llm.base_utils import BaseTokenCounter
|
||||
from litellm.proxy.proxy_server import _try_provider_token_count, token_counter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
|
|
|
|||
|
|
@ -1,21 +1,28 @@
|
|||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
from litellm.llms.gemini.count_tokens.transformation import (
|
||||
GeminiCountTokensPayload,
|
||||
InvalidAnthropicRequest,
|
||||
build_count_tokens_payload,
|
||||
)
|
||||
|
||||
MODEL = "gemini-2.5-flash"
|
||||
|
||||
|
||||
def _payload(messages, system=None, tools=None) -> GeminiCountTokensPayload:
|
||||
payload = build_count_tokens_payload({"model": MODEL, "messages": messages, "system": system, "tools": tools})
|
||||
assert isinstance(payload, GeminiCountTokensPayload), payload
|
||||
return payload
|
||||
|
||||
|
||||
def test_anthropic_tool_turns_become_gemini_function_call_and_response_parts():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[
|
||||
payload = _payload(
|
||||
[
|
||||
{"role": "user", "content": "What is the weather in Paris?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {"city": "Paris"}}],
|
||||
},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "Sunny"}]},
|
||||
],
|
||||
system=None,
|
||||
tools=None,
|
||||
]
|
||||
)
|
||||
|
||||
assert payload.contents == [
|
||||
|
|
@ -30,12 +37,37 @@ def test_anthropic_tool_turns_become_gemini_function_call_and_response_parts():
|
|||
assert payload.tools is None
|
||||
|
||||
|
||||
def test_tool_result_with_block_list_content_keeps_its_text():
|
||||
payload = _payload(
|
||||
[
|
||||
{"role": "user", "content": "Weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_1",
|
||||
"content": [{"type": "text", "text": "Sunny"}, {"type": "text", "text": "21C"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert payload.contents[2] == {
|
||||
"role": "user",
|
||||
"parts": [{"function_response": {"name": "get_weather", "response": {"content": "Sunny21C"}}}],
|
||||
}, payload.contents
|
||||
|
||||
|
||||
def test_system_prompt_is_lifted_out_of_contents_into_system_instruction():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
payload = _payload(
|
||||
[{"role": "user", "content": "hi"}],
|
||||
system=[{"type": "text", "text": "Be terse"}, {"type": "text", "text": "Answer in French"}],
|
||||
tools=None,
|
||||
)
|
||||
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
|
|
@ -45,19 +77,15 @@ def test_system_prompt_is_lifted_out_of_contents_into_system_instruction():
|
|||
|
||||
|
||||
def test_string_system_prompt_becomes_system_instruction():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL, messages=[{"role": "user", "content": "hi"}], system="You are terse", tools=None
|
||||
)
|
||||
payload = _payload([{"role": "user", "content": "hi"}], system="You are terse")
|
||||
|
||||
assert payload.system_instruction == {"parts": [{"text": "You are terse"}]}
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
|
||||
|
||||
def test_anthropic_tools_become_gemini_function_declarations():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system=None,
|
||||
payload = _payload(
|
||||
[{"role": "user", "content": "hi"}],
|
||||
tools=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
|
|
@ -89,10 +117,24 @@ def test_anthropic_tools_become_gemini_function_declarations():
|
|||
], payload.tools
|
||||
|
||||
|
||||
def test_unrecognised_system_value_is_dropped_not_sent():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL, messages=[{"role": "user", "content": "hi"}], system={"unexpected": "shape"}, tools=None
|
||||
def test_system_of_an_unrecognised_shape_is_rejected_as_invalid_not_sent():
|
||||
result = build_count_tokens_payload(
|
||||
{
|
||||
"model": MODEL,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"system": {"unexpected": "shape"},
|
||||
"tools": None,
|
||||
}
|
||||
)
|
||||
|
||||
assert payload.system_instruction is None
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
assert isinstance(result, InvalidAnthropicRequest), result
|
||||
assert "system" in result.message, result.message
|
||||
|
||||
|
||||
def test_tools_that_are_not_a_list_are_rejected_as_invalid():
|
||||
result = build_count_tokens_payload(
|
||||
{"model": MODEL, "messages": [{"role": "user", "content": "hi"}], "system": None, "tools": {"name": "one"}}
|
||||
)
|
||||
|
||||
assert isinstance(result, InvalidAnthropicRequest), result
|
||||
assert "tools" in result.message, result.message
|
||||
|
|
|
|||
|
|
@ -40,18 +40,10 @@ class TestGeminiModelInfo:
|
|||
# Test edge cases where model names end with characters from "models/"
|
||||
# These would be incorrectly processed if using strip("models/") instead of replace("models/", "")
|
||||
models = [
|
||||
{
|
||||
"name": "models/gemini-1.5-pro"
|
||||
}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
|
||||
{
|
||||
"name": "models/test-model"
|
||||
}, # ends with 'l' - would become "gemini/test-mode" with strip()
|
||||
{
|
||||
"name": "models/custom-models"
|
||||
}, # ends with 's' - would become "gemini/custom-model" with strip()
|
||||
{
|
||||
"name": "models/demo"
|
||||
}, # ends with 'o' - would become "gemini/dem" with strip()
|
||||
{"name": "models/gemini-1.5-pro"}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
|
||||
{"name": "models/test-model"}, # ends with 'l' - would become "gemini/test-mode" with strip()
|
||||
{"name": "models/custom-models"}, # ends with 's' - would become "gemini/custom-model" with strip()
|
||||
{"name": "models/demo"}, # ends with 'o' - would become "gemini/dem" with strip()
|
||||
]
|
||||
|
||||
result = gemini_model_info.process_model_name(models)
|
||||
|
|
@ -102,16 +94,10 @@ class TestGoogleAIStudioTokenCounter:
|
|||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
# Test with gemini provider - should return True
|
||||
assert (
|
||||
token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value)
|
||||
is True
|
||||
)
|
||||
assert token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value) is True
|
||||
|
||||
# Test with other providers - should return False
|
||||
assert (
|
||||
token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value)
|
||||
is False
|
||||
)
|
||||
assert token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value) is False
|
||||
assert token_counter.should_use_token_counting_api("anthropic") is False
|
||||
assert token_counter.should_use_token_counting_api("vertex_ai") is False
|
||||
|
||||
|
|
@ -162,11 +148,19 @@ class TestGoogleAIStudioTokenCounter:
|
|||
|
||||
# Verify the mock was called correctly
|
||||
mock_acount_tokens.assert_called_once_with(
|
||||
model=model_to_use, contents=contents, client=None
|
||||
model=model_to_use,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
contents=contents,
|
||||
system_instruction=None,
|
||||
tools=None,
|
||||
client=None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _counter_with_upstream(upstream_response: httpx.Response) -> tuple[GoogleAIStudioTokenCounter, list[httpx.Request]]:
|
||||
def _counter_with_upstream(
|
||||
upstream_response: httpx.Response,
|
||||
) -> tuple[GoogleAIStudioTokenCounter, list[httpx.Request]]:
|
||||
seen_requests: list[httpx.Request] = []
|
||||
|
||||
def upstream(request: httpx.Request) -> httpx.Response:
|
||||
|
|
@ -186,7 +180,9 @@ class TestGoogleAIStudioTokenCounter:
|
|||
contents=None,
|
||||
deployment={"litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
tools=[{"name": "get_weather", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}}],
|
||||
tools=[
|
||||
{"name": "get_weather", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}}
|
||||
],
|
||||
system="Be terse",
|
||||
)
|
||||
|
||||
|
|
@ -199,7 +195,14 @@ class TestGoogleAIStudioTokenCounter:
|
|||
"contents": [{"role": "user", "parts": [{"text": "What is the weather in Paris?"}]}],
|
||||
"systemInstruction": {"parts": [{"text": "Be terse"}]},
|
||||
"tools": [
|
||||
{"function_declarations": [{"name": "get_weather", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}}]}
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
}
|
||||
}
|
||||
|
|
@ -211,6 +214,27 @@ class TestGoogleAIStudioTokenCounter:
|
|||
original_response={"totalTokens": 42},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_system_without_tools_still_wraps_in_generate_content_request(self):
|
||||
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 9}))
|
||||
|
||||
await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
system=[{"type": "text", "text": "Be terse"}],
|
||||
)
|
||||
|
||||
assert json.loads(seen[0].content) == {
|
||||
"generateContentRequest": {
|
||||
"model": "models/gemini-2.5-flash",
|
||||
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
|
||||
"systemInstruction": {"parts": [{"text": "Be terse"}]},
|
||||
}
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_contents_are_sent_unchanged(self):
|
||||
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 3}))
|
||||
|
|
@ -299,9 +323,7 @@ class TestGoogleAIStudioTokenCounter:
|
|||
"functionResponse": {
|
||||
"id": "read_many_files-1757526647518-730a691aac11c", # This should be removed
|
||||
"name": "read_many_files",
|
||||
"response": {
|
||||
"output": "No files matching the criteria were found or all were skipped."
|
||||
},
|
||||
"response": {"output": "No files matching the criteria were found or all were skipped."},
|
||||
}
|
||||
}
|
||||
],
|
||||
|
|
@ -310,9 +332,7 @@ class TestGoogleAIStudioTokenCounter:
|
|||
]
|
||||
|
||||
# Clean the contents
|
||||
cleaned_contents = token_counter._clean_contents_for_gemini_api(
|
||||
contents_with_id
|
||||
)
|
||||
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_with_id)
|
||||
|
||||
# Verify the 'id' field was removed
|
||||
function_response = cleaned_contents[1]["parts"][0]["functionResponse"]
|
||||
|
|
@ -321,8 +341,7 @@ class TestGoogleAIStudioTokenCounter:
|
|||
assert "response" in function_response
|
||||
assert function_response["name"] == "read_many_files"
|
||||
assert (
|
||||
function_response["response"]["output"]
|
||||
== "No files matching the criteria were found or all were skipped."
|
||||
function_response["response"]["output"] == "No files matching the criteria were found or all were skipped."
|
||||
)
|
||||
|
||||
def test_clean_contents_for_gemini_api_preserves_other_fields(self):
|
||||
|
|
@ -338,9 +357,7 @@ class TestGoogleAIStudioTokenCounter:
|
|||
]
|
||||
|
||||
# Clean the contents
|
||||
cleaned_contents = token_counter._clean_contents_for_gemini_api(
|
||||
contents_without_function_response
|
||||
)
|
||||
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_without_function_response)
|
||||
|
||||
# Verify the contents are unchanged
|
||||
assert cleaned_contents == contents_without_function_response
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue