mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(vertex_ai): forward system and tools to partner model count_tokens (#43900)
* fix(vertex_ai): forward system and tools to partner model count_tokens Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vertex_ai): avoid mutable token request construction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vertex_ai): return partner count_tokens provider errors as values so the proxy falls back locally * test(integration): cover Vertex AI partner count_tokens forwarding and local fallbacks Add wire-level cells for /v1/messages/count_tokens, /utils/token_counter, /v1/responses/input_tokens and the Gemini countTokens route on a Vertex AI Claude deployment: the system prompt and tools reach the partner count-tokens endpoint verbatim, null fields stay out of the body, malformed tools are rejected before any peer call, peer, token-endpoint and connection failures fall back to the local tokenizer unless disable_token_counter is set, generation on the same deployment keeps working, and concurrent bursts survive a peer outage, a slow peer and a worker SIGKILL. The sdk cells cover litellm.acount_tokens the same way. The _support/process.py and _support/client.py harness files are brought to main's content so the self-booting cells read INTEGRATION_PROXY_READY_SECONDS instead of a fixed 70 s boot budget. --------- Co-authored-by: jesus <jesus@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
1ef0fe9790
commit
6d73fa6b49
6 changed files with 1027 additions and 8 deletions
|
|
@ -1240,6 +1240,9 @@ class VertexAITokenCounter(BaseTokenCounter):
|
|||
) -> TokenCountResponse | None:
|
||||
import copy
|
||||
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
|
||||
VertexAIError as PartnerVertexAIError,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
|
||||
VertexAIPartnerModels,
|
||||
)
|
||||
|
|
@ -1269,14 +1272,32 @@ class VertexAITokenCounter(BaseTokenCounter):
|
|||
"vertex_ai_credentials"
|
||||
)
|
||||
|
||||
result = await partner_models_handler.count_tokens(
|
||||
model=model_to_use,
|
||||
messages=messages or [],
|
||||
litellm_params=partner_litellm_params,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_credentials=vertex_credentials,
|
||||
)
|
||||
try:
|
||||
result = await partner_models_handler.count_tokens(
|
||||
model=model_to_use,
|
||||
messages=messages or [],
|
||||
litellm_params=partner_litellm_params,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_credentials=vertex_credentials,
|
||||
system=system,
|
||||
tools=tools,
|
||||
)
|
||||
except (PartnerVertexAIError, httpx.HTTPStatusError) as e:
|
||||
status_code: Final = e.response.status_code
|
||||
error_message: Final = e.message if isinstance(e, PartnerVertexAIError) else e.response.text
|
||||
verbose_logger.warning(
|
||||
"Vertex AI partner CountTokens API error: status=%s, message=%s", status_code, error_message
|
||||
)
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="vertex_ai_partner_models",
|
||||
error=True,
|
||||
error_message=error_message,
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
if result is not None:
|
||||
return TokenCountResponse(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
## API Handler for calling Vertex AI Partner Models
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -263,6 +264,8 @@ class VertexAIPartnerModels(VertexBase):
|
|||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_credentials=None,
|
||||
system: object | None = None,
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
):
|
||||
"""
|
||||
Count tokens for Vertex AI partner models (Anthropic Claude, Mistral, etc.)
|
||||
|
|
@ -296,6 +299,9 @@ class VertexAIPartnerModels(VertexBase):
|
|||
request_data: Final = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**MappingProxyType(
|
||||
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
|
||||
),
|
||||
}
|
||||
|
||||
# Prepare litellm_params with credentials
|
||||
|
|
|
|||
31
tests/integration/_support/vertex.py
Normal file
31
tests/integration/_support/vertex.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
|
||||
|
||||
def service_account_json(project: str, token_url: str) -> str:
|
||||
private_key: Final = (
|
||||
rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
)
|
||||
.decode()
|
||||
)
|
||||
return json.dumps(
|
||||
{
|
||||
"type": "service_account",
|
||||
"project_id": project,
|
||||
"private_key_id": "scripted",
|
||||
"private_key": private_key,
|
||||
"client_email": f"scripted@{project}.iam.gserviceaccount.com",
|
||||
"client_id": "0",
|
||||
"auth_uri": f"{token_url}/_oauth/authorize",
|
||||
"token_uri": f"{token_url}/_oauth/token",
|
||||
}
|
||||
)
|
||||
|
|
@ -0,0 +1,728 @@
|
|||
import json
|
||||
import re
|
||||
import signal
|
||||
import socket
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import ExitStack
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, eventually
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.vertex import service_account_json
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "claude-sonnet-4-6"
|
||||
_PROJECT: Final = "scripted-project"
|
||||
_LOCATION: Final = "us-east5"
|
||||
_MODELS_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/anthropic/models"
|
||||
_COUNT_TARGET: Final = f"{_MODELS_PATH}/count-tokens:rawPredict"
|
||||
_MESSAGE_TARGET: Final = f"{_MODELS_PATH}/{_BACKEND}:rawPredict"
|
||||
_STREAM_TARGET: Final = f"{_MODELS_PATH}/{_BACKEND}:streamRawPredict?alt=sse"
|
||||
_PEER_COUNT: Final = 4242
|
||||
_REJECTION: Final = "scripted partner rejection"
|
||||
_REJECT_TEXT: Final = "The peer must reject this message"
|
||||
_REPLY_TEXT: Final = "scripted reply"
|
||||
_OWNED_MODEL: Final = "partner-claude"
|
||||
_OWNED_UNREACHABLE_MODEL: Final = "partner-claude-unreachable"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
|
||||
_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this message"}]
|
||||
_SYSTEM: Final = "You are a terse assistant that answers in one sentence"
|
||||
_SYSTEM_BLOCKS: Final[list[JsonValue]] = [
|
||||
{"type": "text", "text": "You are a terse assistant"},
|
||||
{"type": "text", "text": "Answer in one sentence"},
|
||||
]
|
||||
_WEATHER_SCHEMA: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string", "description": "City to look up"}},
|
||||
"required": ["city"],
|
||||
}
|
||||
_TOOLS: Final[list[JsonValue]] = [
|
||||
{"name": "get_weather", "description": "Look up the current weather for a city", "input_schema": _WEATHER_SCHEMA}
|
||||
]
|
||||
_OPENAI_TOOLS: Final[list[JsonValue]] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Look up the current weather for a city",
|
||||
"parameters": _WEATHER_SCHEMA,
|
||||
},
|
||||
}
|
||||
]
|
||||
_RESPONSES_TOOLS: Final[list[JsonValue]] = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Look up the current weather for a city",
|
||||
"parameters": _WEATHER_SCHEMA,
|
||||
}
|
||||
]
|
||||
_PEER_BARE: Final[dict[str, JsonValue]] = {"model": _BACKEND, "messages": _MESSAGES}
|
||||
_PEER_FULL: Final[dict[str, JsonValue]] = {**_PEER_BARE, "system": _SYSTEM, "tools": _TOOLS}
|
||||
_GEMINI_BODY: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "Count this"}]}]}
|
||||
_GEMINI_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this"}]
|
||||
|
||||
_REPLY: Final[dict[str, JsonValue]] = {
|
||||
"id": "msg_scripted",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": _BACKEND,
|
||||
"content": [{"type": "text", "text": _REPLY_TEXT}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 3},
|
||||
}
|
||||
_EVENTS: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = (
|
||||
(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {**_REPLY, "content": [], "stop_reason": None, "usage": {"input_tokens": 5, "output_tokens": 1}},
|
||||
},
|
||||
),
|
||||
("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
|
||||
(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _REPLY_TEXT}},
|
||||
),
|
||||
("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {"output_tokens": 3},
|
||||
},
|
||||
),
|
||||
("message_stop", {"type": "message_stop"}),
|
||||
)
|
||||
_SSE: Final = tuple(f"event: {name}\ndata: {json.dumps(data)}\n\n".encode() for name, data in _EVENTS)
|
||||
|
||||
|
||||
def _counted(_request: Request) -> Reply:
|
||||
return Reply(body=json.dumps({"input_tokens": _PEER_COUNT}).encode())
|
||||
|
||||
|
||||
def _rejected(status: int) -> Reply:
|
||||
return Reply(
|
||||
status=status,
|
||||
body=json.dumps({"type": "error", "error": {"type": "invalid_request_error", "message": _REJECTION}}).encode(),
|
||||
)
|
||||
|
||||
|
||||
def _rejecting(status: int) -> Callable[[Request], Reply]:
|
||||
def count(_request: Request) -> Reply:
|
||||
return _rejected(status)
|
||||
|
||||
return count
|
||||
|
||||
|
||||
def _anthropic_message(message: JsonValue) -> bool:
|
||||
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
|
||||
|
||||
|
||||
def _anthropic_tool(tool: JsonValue) -> bool:
|
||||
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
|
||||
|
||||
|
||||
def _strict(request: Request) -> Reply:
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
messages: Final = body.get("messages")
|
||||
tools: Final = body.get("tools", [])
|
||||
accepted: Final = (
|
||||
isinstance(messages, list)
|
||||
and all(map(_anthropic_message, messages))
|
||||
and isinstance(body.get("system", ""), (str, list))
|
||||
and isinstance(tools, list)
|
||||
and all(map(_anthropic_tool, tools))
|
||||
)
|
||||
return _counted(request) if accepted else _rejected(400)
|
||||
|
||||
|
||||
def _rejecting_marked_messages(request: Request) -> Reply:
|
||||
return _rejected(400) if _REJECT_TEXT in request.body.decode() else _counted(request)
|
||||
|
||||
|
||||
def _peer(count: Callable[[Request], Reply] = _counted) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.target.endswith("/count-tokens:rawPredict"):
|
||||
return count(request)
|
||||
if request.target.endswith(":streamRawPredict?alt=sse"):
|
||||
return Reply(content_type="text/event-stream", chunks=_SSE)
|
||||
if request.target.endswith(f"/{_BACKEND}:rawPredict"):
|
||||
return Reply(body=json.dumps(_REPLY).encode())
|
||||
return Reply(status=404, body=json.dumps({"error": f"unscripted target {request.target}"}).encode())
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _count_requests(requests: Sequence[Request]) -> tuple[Request, ...]:
|
||||
return tuple(request for request in requests if "count-tokens" in request.target)
|
||||
|
||||
|
||||
def _count_bodies(requests: Sequence[Request], target: str = _COUNT_TARGET) -> tuple[dict[str, JsonValue], ...]:
|
||||
counts: Final = _count_requests(requests)
|
||||
for request in counts:
|
||||
assert (request.method, request.target) == ("POST", target), request.target
|
||||
assert request.headers["authorization"] == "Bearer scripted-token", request.headers
|
||||
return tuple(_JSON_OBJECT.validate_json(request.body) for request in counts)
|
||||
|
||||
|
||||
def _counted_bodies(wire: Wire) -> tuple[dict[str, JsonValue], ...]:
|
||||
return _count_bodies(wire.drain())
|
||||
|
||||
|
||||
def _bare(model: str) -> dict[str, JsonValue]:
|
||||
return {"model": model, "messages": _MESSAGES}
|
||||
|
||||
|
||||
def _full(model: str) -> dict[str, JsonValue]:
|
||||
return {**_bare(model), "system": _SYSTEM, "tools": _TOOLS}
|
||||
|
||||
|
||||
def _deployment(gateway: Gateway, scenario: Scenario, api_base: str, **overrides: JsonValue) -> str:
|
||||
return scenario.model(
|
||||
**{
|
||||
"model": f"vertex_ai/{_BACKEND}",
|
||||
"api_base": api_base,
|
||||
"api_key": None,
|
||||
"vertex_project": _PROJECT,
|
||||
"vertex_location": _LOCATION,
|
||||
"vertex_credentials": service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")),
|
||||
**overrides,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _count(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response:
|
||||
return gateway.request("POST", "/v1/messages/count_tokens", body)
|
||||
|
||||
|
||||
def _payload(response: httpx.Response) -> dict[str, JsonValue]:
|
||||
assert response.status_code == 200, response.text
|
||||
return _JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int:
|
||||
response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"})
|
||||
payload: Final = _payload(response)
|
||||
total: Final = payload["total_tokens"]
|
||||
assert payload["tokenizer_type"] != "vertex_ai_partner_models", response.text
|
||||
assert isinstance(total, int) and 0 < total != _PEER_COUNT, response.text
|
||||
return total
|
||||
|
||||
|
||||
def _closed_port_url() -> str:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
return f"http://127.0.0.1:{reserve.getsockname()[1]}"
|
||||
|
||||
|
||||
def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]:
|
||||
return tuple(
|
||||
stack.enter_context(httpx.Client(base_url=base_url, timeout=30, trust_env=False)) for _ in range(count)
|
||||
)
|
||||
|
||||
|
||||
def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]:
|
||||
response: Final = client.post(
|
||||
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"}
|
||||
)
|
||||
return response.status_code, _JSON_OBJECT.validate_json(response.content).get("input_tokens")
|
||||
|
||||
|
||||
def _generated_then_counted(
|
||||
client: httpx.Client, key: str, model: str, body: Mapping[str, JsonValue]
|
||||
) -> tuple[int, int, JsonValue]:
|
||||
generated: Final = client.post(
|
||||
"/v1/messages",
|
||||
json={
|
||||
"model": model,
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": f"Generate before counting {uuid.uuid4().hex}"}],
|
||||
},
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
return generated.status_code, *_counted_on(client, key, body)
|
||||
|
||||
|
||||
def _local_port(client: httpx.Client) -> int:
|
||||
with client.stream("GET", "/health/liveliness") as response:
|
||||
port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1])
|
||||
response.read()
|
||||
assert response.status_code == 200, response.text
|
||||
return port
|
||||
|
||||
|
||||
def _counted_or_dropped(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue] | None:
|
||||
try:
|
||||
return _counted_on(client, key, body)
|
||||
except httpx.TransportError:
|
||||
return None
|
||||
|
||||
|
||||
def _accepted_client_ports(pid: int, proxy_port: int) -> frozenset[int]:
|
||||
return frozenset(
|
||||
connection.raddr.port
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.raddr and connection.laddr.port == proxy_port
|
||||
)
|
||||
|
||||
|
||||
def _owned_config(
|
||||
path: Path, gateway: Gateway, api_bases: Mapping[str, str], settings: Mapping[str, JsonValue]
|
||||
) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
**config,
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {
|
||||
"model": f"vertex_ai/{_BACKEND}",
|
||||
"api_base": api_base,
|
||||
"vertex_project": _PROJECT,
|
||||
"vertex_location": _LOCATION,
|
||||
"vertex_credentials": service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")),
|
||||
},
|
||||
}
|
||||
for name, api_base in api_bases.items()
|
||||
],
|
||||
"litellm_settings": {**config["litellm_settings"], **settings},
|
||||
}
|
||||
)
|
||||
)
|
||||
return path
|
||||
|
||||
|
||||
@pytest.mark.parametrize("system", [_SYSTEM, _SYSTEM_BLOCKS], ids=["string", "blocks"])
|
||||
def test_messages_count_tokens_forwards_system_and_tools(gateway: Gateway, system: JsonValue) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "system": system, "tools": _TOOLS})
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "system": system, "tools": _TOOLS},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_without_system_or_tools_sends_bare_body(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, _bare(model))
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == (_PEER_BARE,)
|
||||
|
||||
|
||||
def test_utils_token_counter_call_endpoint_forwards_system_and_tools(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/utils/token_counter", _full(model), params={"call_endpoint": "true"}
|
||||
)
|
||||
payload: Final = _payload(response)
|
||||
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_PEER_COUNT, "vertex_ai_partner_models")
|
||||
assert (payload["request_model"], payload["model_used"]) == (model, _BACKEND), response.text
|
||||
assert _counted_bodies(wire) == (_PEER_FULL,)
|
||||
|
||||
|
||||
def test_utils_token_counter_local_mode_never_calls_the_peer(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
assert _local_count(gateway, _full(model)) > 0
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
def test_utils_token_counter_falls_back_locally_when_peer_rejects_openai_tools(gateway: Gateway) -> None:
|
||||
with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
body: Final = {**_bare(model), "tools": _OPENAI_TOOLS}
|
||||
response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "true"})
|
||||
payload: Final = _payload(response)
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "tools": _OPENAI_TOOLS},)
|
||||
assert payload["total_tokens"] == _local_count(gateway, body), response.text
|
||||
assert payload["tokenizer_type"] != "vertex_ai_partner_models", response.text
|
||||
|
||||
|
||||
def test_responses_input_tokens_counts_through_the_partner_peer(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/responses/input_tokens", {"model": model, "input": "Count this message"}
|
||||
)
|
||||
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == (_PEER_BARE,)
|
||||
|
||||
|
||||
def test_responses_input_tokens_falls_back_locally_when_peer_rejects(gateway: Gateway) -> None:
|
||||
with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses/input_tokens",
|
||||
{"model": model, "input": "Count this message", "instructions": "Be terse", "tools": _RESPONSES_TOOLS},
|
||||
)
|
||||
payload: Final = _payload(response)
|
||||
(sent,) = _counted_bodies(wire)
|
||||
assert sent["tools"] == _RESPONSES_TOOLS, sent
|
||||
local: Final = _local_count(gateway, {"model": model, "messages": sent["messages"], "tools": _RESPONSES_TOOLS})
|
||||
assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text
|
||||
|
||||
|
||||
def test_gemini_count_tokens_route_reaches_the_partner_peer(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY)
|
||||
assert "totalTokens" in _payload(response), response.text
|
||||
assert _counted_bodies(wire) == ({"model": _BACKEND, "messages": _GEMINI_MESSAGES},)
|
||||
|
||||
|
||||
def test_gemini_count_tokens_route_falls_back_locally_when_peer_rejects(gateway: Gateway) -> None:
|
||||
with wire_server(_peer(_rejecting(400))) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY)
|
||||
payload: Final = _payload(response)
|
||||
assert _counted_bodies(wire) == ({"model": _BACKEND, "messages": _GEMINI_MESSAGES},)
|
||||
assert payload["totalTokens"] == _local_count(gateway, {"model": model, "messages": _GEMINI_MESSAGES})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [400, 500, 503])
|
||||
def test_messages_count_tokens_falls_back_locally_when_peer_errors(gateway: Gateway, status: int) -> None:
|
||||
with wire_server(_peer(_rejecting(status))) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, _full(model))
|
||||
payload: Final = _payload(response)
|
||||
assert _counted_bodies(wire) == (_PEER_FULL,)
|
||||
assert payload == {"input_tokens": _local_count(gateway, _full(model))}, response.text
|
||||
|
||||
|
||||
def test_messages_count_tokens_falls_back_locally_when_token_endpoint_rejects(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.target == "/_oauth/token":
|
||||
return Reply(
|
||||
status=400,
|
||||
body=json.dumps({"error": "invalid_grant", "error_description": "scripted refusal"}).encode(),
|
||||
)
|
||||
return _peer()(request)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(
|
||||
gateway, scenario, wire.url, vertex_credentials=service_account_json(_PROJECT, wire.url)
|
||||
)
|
||||
response: Final = _count(gateway, _full(model))
|
||||
payload: Final = _payload(response)
|
||||
targets: Final = frozenset(request.target for request in wire.drain())
|
||||
assert targets == {"/_oauth/token"}, targets
|
||||
assert payload == {"input_tokens": _local_count(gateway, _full(model))}, response.text
|
||||
|
||||
|
||||
def test_messages_count_tokens_falls_back_locally_when_peer_is_unreachable(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, _closed_port_url())
|
||||
response: Final = _count(gateway, _full(model))
|
||||
assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tools",
|
||||
[5, "", "x" * 5120, ["get_weather"]],
|
||||
ids=["int", "empty_string", "5kb_string", "list_of_strings"],
|
||||
)
|
||||
def test_messages_count_tokens_rejects_malformed_tools_without_calling_the_peer(
|
||||
gateway: Gateway, tools: JsonValue
|
||||
) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
refused: Final = _count(gateway, {**_bare(model), "tools": tools})
|
||||
assert 400 <= refused.status_code < 600, refused.text
|
||||
assert "input_tokens" not in refused.text, refused.text
|
||||
assert wire.drain() == ()
|
||||
assert _generated_then_counted(gateway.client, gateway.key, model, _bare(model)) == (200, 200, _PEER_COUNT)
|
||||
assert _counted_bodies(wire) == (_PEER_BARE,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_forwards_an_empty_tools_list(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "tools": []})
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "tools": []},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_forwards_an_empty_system_string(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "system": ""})
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "system": ""},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_leaves_null_system_and_tools_out(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "system": None, "tools": None})
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == (_PEER_BARE,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_falls_back_locally_when_peer_rejects_a_non_text_system(gateway: Gateway) -> None:
|
||||
with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "system": 5})
|
||||
payload: Final = _payload(response)
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "system": 5},)
|
||||
assert payload == {"input_tokens": _local_count(gateway, _bare(model))}, response.text
|
||||
|
||||
|
||||
def test_messages_count_tokens_forwards_a_5kb_system_verbatim(gateway: Gateway) -> None:
|
||||
system: Final = "Answer in one sentence. " * 214
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {**_bare(model), "system": system})
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
assert _counted_bodies(wire) == ({**_PEER_BARE, "system": system},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_duplicate_system_and_tools_keys_forward_one_value_each(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
fields: Final = f'"system": {json.dumps(_SYSTEM)}, "tools": {json.dumps(_TOOLS)}'
|
||||
response: Final = gateway.client.post(
|
||||
"/v1/messages/count_tokens",
|
||||
content=f'{{"model": "{model}", "messages": {json.dumps(_MESSAGES)}, {fields}, {fields}}}',
|
||||
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
|
||||
)
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
(sent,) = _count_requests(wire.drain())
|
||||
assert _count_bodies((sent,)) == (_PEER_FULL,)
|
||||
assert (sent.body.count(b'"system"'), sent.body.count(b'"tools"')) == (1, 1), sent.body
|
||||
|
||||
|
||||
def test_messages_count_tokens_unauthenticated_request_never_reaches_the_peer(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/messages/count_tokens", _full(model), key="sk-not-a-key-this-proxy-issued"
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert "input_tokens" not in response.text, response.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fields", [{}, {"messages": []}], ids=["missing", "empty"])
|
||||
def test_messages_count_tokens_without_messages_is_rejected(gateway: Gateway, fields: dict[str, JsonValue]) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = _count(gateway, {"model": model, "system": _SYSTEM, "tools": _TOOLS, **fields})
|
||||
assert response.status_code == 400, response.text
|
||||
assert "messages parameter is required" in response.text, response.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
def test_count_tokens_location_override_targets_the_count_region(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(
|
||||
gateway, scenario, wire.url, vertex_location="global", vertex_count_tokens_location="europe-west1"
|
||||
)
|
||||
response: Final = _count(gateway, _bare(model))
|
||||
assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text
|
||||
target: Final = _COUNT_TARGET.replace(f"/locations/{_LOCATION}/", "/locations/europe-west1/")
|
||||
assert _count_bodies(wire.drain(), target) == (_PEER_BARE,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_repeated_request_reaches_the_peer_each_time(gateway: Gateway) -> None:
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
answers: Final = tuple(_payload(_count(gateway, _bare(model))) for _ in range(2))
|
||||
assert answers == ({"input_tokens": _PEER_COUNT},) * 2
|
||||
assert _counted_bodies(wire) == (_PEER_BARE,) * 2
|
||||
|
||||
|
||||
@pytest.mark.timeout(240) # boots an owned two-worker proxy with litellm_settings.disable_token_counter
|
||||
def test_disabled_token_counter_surfaces_provider_failures_instead_of_counting_locally(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
with wire_server(_peer(_rejecting_marked_messages)) as wire:
|
||||
config: Final = _owned_config(
|
||||
tmp_path / "disabled-token-counter.yaml",
|
||||
gateway,
|
||||
{_OWNED_MODEL: wire.url, _OWNED_UNREACHABLE_MODEL: _closed_port_url()},
|
||||
{"disable_token_counter": True},
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
counted: Final = _count(owned.gateway, _full(_OWNED_MODEL))
|
||||
assert _payload(counted) == {"input_tokens": _PEER_COUNT}, counted.text
|
||||
rejected: Final = _count(
|
||||
owned.gateway, {**_full(_OWNED_MODEL), "messages": [{"role": "user", "content": _REJECT_TEXT}]}
|
||||
)
|
||||
assert rejected.status_code == 400, rejected.text
|
||||
assert _REJECTION in rejected.text and "input_tokens" not in rejected.text, rejected.text
|
||||
unreachable: Final = _count(owned.gateway, _full(_OWNED_UNREACHABLE_MODEL))
|
||||
assert 500 <= unreachable.status_code < 600, unreachable.text
|
||||
assert "input_tokens" not in unreachable.text, unreachable.text
|
||||
local: Final = owned.gateway.request(
|
||||
"POST", "/utils/token_counter", _full(_OWNED_MODEL), params={"call_endpoint": "false"}
|
||||
)
|
||||
assert local.status_code == 503, local.text
|
||||
assert len(_counted_bodies(wire)) == 2
|
||||
|
||||
|
||||
def test_peer_outage_between_concurrent_waves_falls_back_then_recovers(gateway: Gateway) -> None:
|
||||
with ExitStack() as stack:
|
||||
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
scenario: Final = stack.enter_context(gateway.scenario())
|
||||
with wire_server(_peer()) as wire:
|
||||
port: Final = int(wire.url.rsplit(":", 1)[1])
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
body: Final = _full(model)
|
||||
local: Final = _local_count(gateway, body)
|
||||
|
||||
def generate_then_count(client: httpx.Client) -> tuple[int, int, JsonValue]:
|
||||
return _generated_then_counted(client, gateway.key, model, body)
|
||||
|
||||
assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _PEER_COUNT),) * len(clients)
|
||||
assert _counted_bodies(wire) == (_PEER_FULL,) * len(clients)
|
||||
outage: Final = tuple(pool.map(lambda client: _counted_on(client, gateway.key, body), clients))
|
||||
assert outage == ((200, local),) * len(clients)
|
||||
with wire_server(_peer(), port=port) as revived:
|
||||
assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _PEER_COUNT),) * len(clients)
|
||||
assert _counted_bodies(revived) == (_PEER_FULL,) * len(clients)
|
||||
|
||||
|
||||
def test_slow_peer_holds_concurrent_counts_without_stalling_the_proxy(gateway: Gateway) -> None:
|
||||
held: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
release: Final = threading.Event()
|
||||
|
||||
def hold(request: Request) -> Reply:
|
||||
held.put(request.target)
|
||||
assert release.wait(timeout=20), "Held count was never released"
|
||||
return _counted(request)
|
||||
|
||||
with ExitStack() as stack:
|
||||
clients: Final = _clients(stack, str(gateway.client.base_url), 6)
|
||||
wire: Final = stack.enter_context(wire_server(_peer(hold)))
|
||||
scenario: Final = stack.enter_context(gateway.scenario())
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
try:
|
||||
futures: Final = tuple(
|
||||
pool.submit(_generated_then_counted, client, gateway.key, model, _full(model)) for client in clients
|
||||
)
|
||||
eventually(held.qsize, lambda size: size == len(clients), seconds=30)
|
||||
assert gateway.request("GET", "/health/liveliness").status_code == 200
|
||||
assert _local_count(gateway, _full(model)) > 0
|
||||
assert not any(future.done() for future in futures)
|
||||
finally:
|
||||
release.set()
|
||||
assert tuple(future.result(timeout=30) for future in futures) == ((200, 200, _PEER_COUNT),) * len(clients)
|
||||
assert len(_counted_bodies(wire)) == len(clients)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300) # boots an owned two-worker proxy, kills one worker, and waits for its replacement
|
||||
def test_worker_sigkill_mid_burst_leaves_the_sibling_counting(gateway: Gateway, tmp_path: Path) -> None:
|
||||
held: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
release: Final = threading.Event()
|
||||
|
||||
def hold(request: Request) -> Reply:
|
||||
held.put(request.target)
|
||||
assert release.wait(timeout=60), "Held count was never released"
|
||||
return _counted(request)
|
||||
|
||||
with ExitStack() as stack:
|
||||
wire: Final = stack.enter_context(wire_server(_peer(hold)))
|
||||
stack.callback(release.set)
|
||||
config: Final = _owned_config(tmp_path / "worker-kill.yaml", gateway, {_OWNED_MODEL: wire.url}, {})
|
||||
owned: Final = stack.enter_context(owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2))
|
||||
proxy_url: Final = owned.gateway.client.base_url
|
||||
workers: Final = eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
clients: Final = _clients(stack, str(proxy_url), 12)
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
stack.callback(release.set)
|
||||
ports: Final = tuple(_local_port(client) for client in clients)
|
||||
futures: Final = tuple(
|
||||
pool.submit(_counted_or_dropped, client, gateway.key, _full(_OWNED_MODEL)) for client in clients
|
||||
)
|
||||
eventually(held.qsize, lambda size: size == len(clients), seconds=30)
|
||||
shares: Final = {pid: _accepted_client_ports(pid, proxy_url.port or 0) & frozenset(ports) for pid in workers}
|
||||
assert sum(map(len, shares.values())) == len(clients), shares
|
||||
victim: Final = min((pid for pid in workers if shares[pid]), key=lambda pid: len(shares[pid]))
|
||||
psutil.Process(victim).send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
results: Final = tuple(future.result(timeout=60) for future in futures)
|
||||
for port, result in zip(ports, results, strict=True):
|
||||
assert result == (None if port in shares[victim] else (200, _PEER_COUNT)), (port, result, shares)
|
||||
second_wave: Final = _clients(stack, str(proxy_url), 6)
|
||||
assert tuple(_counted_on(client, gateway.key, _full(_OWNED_MODEL)) for client in second_wave) == (
|
||||
(200, _PEER_COUNT),
|
||||
) * len(second_wave)
|
||||
assert len(_counted_bodies(wire)) == len(clients) + len(second_wave)
|
||||
eventually(lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda started: started >= 3, 120)
|
||||
assert owned.process.poll() is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"])
|
||||
def test_chat_completions_on_the_same_deployment_still_generate(gateway: Gateway, stream: bool) -> None:
|
||||
marker: Final = f"chat control {uuid.uuid4().hex}"
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
"stream": stream,
|
||||
"cache": {"no-cache": True},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _REPLY_TEXT in response.text, response.text
|
||||
assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text
|
||||
(sent,) = wire.drain()
|
||||
assert sent.target == (_STREAM_TARGET if stream else _MESSAGE_TARGET), sent.target
|
||||
assert marker in sent.body.decode(), sent.body
|
||||
|
||||
|
||||
def test_messages_endpoint_on_the_same_deployment_still_generates(gateway: Gateway) -> None:
|
||||
marker: Final = f"messages control {uuid.uuid4().hex}"
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/messages",
|
||||
{"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _REPLY_TEXT in response.text, response.text
|
||||
(sent,) = wire.drain()
|
||||
assert sent.target == _MESSAGE_TARGET, sent.target
|
||||
assert marker in sent.body.decode(), sent.body
|
||||
|
||||
|
||||
def test_responses_endpoint_on_the_same_deployment_still_generates(gateway: Gateway) -> None:
|
||||
marker: Final = f"responses control {uuid.uuid4().hex}"
|
||||
with wire_server(_peer()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire.url)
|
||||
response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": marker})
|
||||
assert response.status_code == 200, response.text
|
||||
assert _REPLY_TEXT in response.text, response.text
|
||||
(sent,) = wire.drain()
|
||||
assert sent.target == _MESSAGE_TARGET, sent.target
|
||||
assert marker in sent.body.decode(), sent.body
|
||||
|
|
@ -0,0 +1,90 @@
|
|||
import json
|
||||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from integration._support.vertex import service_account_json
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "claude-sonnet-4-6"
|
||||
_PROJECT: Final = "scripted-project"
|
||||
_LOCATION: Final = "us-east5"
|
||||
_COUNT_TARGET: Final = (
|
||||
f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/anthropic/models/count-tokens:rawPredict"
|
||||
)
|
||||
_PEER_COUNT: Final = 4242
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_MESSAGES: Final[list[dict[str, str]]] = [{"role": "user", "content": "Count this message"}]
|
||||
_SYSTEM: Final = "You are a terse assistant that answers in one sentence"
|
||||
_TOOLS: Final[list[dict[str, JsonValue]]] = [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Look up the current weather for a city",
|
||||
"input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def _peer(status: int) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.target == "/_oauth/token":
|
||||
token: Final = {"access_token": "scripted-token", "token_type": "Bearer", "expires_in": 3600}
|
||||
return Reply(body=json.dumps(token).encode())
|
||||
if status == 200:
|
||||
return Reply(body=json.dumps({"input_tokens": _PEER_COUNT}).encode())
|
||||
rejection: Final = {"type": "error", "error": {"type": "invalid_request_error", "message": "scripted"}}
|
||||
return Reply(status=status, body=json.dumps(rejection).encode())
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _count_requests(requests: tuple[Request, ...]) -> tuple[Request, ...]:
|
||||
return tuple(request for request in requests if "count-tokens" in request.target)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vertex_environment(monkeypatch: pytest.MonkeyPatch) -> Callable[[str], None]:
|
||||
def configure(token_url: str) -> None:
|
||||
monkeypatch.setenv("VERTEXAI_PROJECT", _PROJECT)
|
||||
monkeypatch.setenv("VERTEXAI_LOCATION", _LOCATION)
|
||||
monkeypatch.setenv("VERTEXAI_CREDENTIALS", service_account_json(_PROJECT, token_url))
|
||||
|
||||
return configure
|
||||
|
||||
|
||||
async def test_acount_tokens_forwards_system_and_tools_to_the_partner_peer(
|
||||
vertex_environment: Callable[[str], None],
|
||||
) -> None:
|
||||
with wire_server(_peer(200)) as wire:
|
||||
vertex_environment(wire.url)
|
||||
counted: Final = await litellm.acount_tokens(
|
||||
model=f"vertex_ai/{_BACKEND}", messages=_MESSAGES, tools=_TOOLS, system=_SYSTEM, api_base=wire.url
|
||||
)
|
||||
(sent,) = _count_requests(wire.drain())
|
||||
assert (sent.method, sent.target, sent.headers["authorization"]) == (
|
||||
"POST",
|
||||
_COUNT_TARGET,
|
||||
"Bearer scripted-token",
|
||||
)
|
||||
assert _JSON_OBJECT.validate_json(sent.body) == {
|
||||
"model": _BACKEND,
|
||||
"messages": _MESSAGES,
|
||||
"system": _SYSTEM,
|
||||
"tools": _TOOLS,
|
||||
}
|
||||
assert (counted.total_tokens, counted.tokenizer_type) == (_PEER_COUNT, "vertex_ai_partner_models"), counted
|
||||
|
||||
|
||||
async def test_acount_tokens_falls_back_to_the_local_tokenizer_when_the_peer_rejects(
|
||||
vertex_environment: Callable[[str], None],
|
||||
) -> None:
|
||||
with wire_server(_peer(400)) as wire:
|
||||
vertex_environment(wire.url)
|
||||
counted: Final = await litellm.acount_tokens(
|
||||
model=f"vertex_ai/{_BACKEND}", messages=_MESSAGES, tools=_TOOLS, system=_SYSTEM, api_base=wire.url
|
||||
)
|
||||
assert len(_count_requests(wire.drain())) == 1
|
||||
assert counted.tokenizer_type == "local_tokenizer", counted
|
||||
assert counted.total_tokens > 0 and counted.total_tokens != _PEER_COUNT, counted
|
||||
|
|
@ -1229,6 +1229,149 @@ async def test_vertex_ai_token_counter_routes_partner_models():
|
|||
assert result.tokenizer_type == "vertex_ai_partner_models"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_ai_token_counter_forwards_system_and_tools_to_partner_request():
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens import handler
|
||||
|
||||
class FakeResponse:
|
||||
status_code = 200
|
||||
|
||||
def json(self) -> dict[str, int]:
|
||||
return {"input_tokens": 37}
|
||||
|
||||
class FakeHttpClient:
|
||||
posted_bodies: tuple[dict[str, object], ...] = ()
|
||||
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
json: dict[str, object],
|
||||
timeout: float,
|
||||
) -> FakeResponse:
|
||||
self.posted_bodies = (*self.posted_bodies, json)
|
||||
return FakeResponse()
|
||||
|
||||
fake_http_client: Final = FakeHttpClient()
|
||||
counter: Final = VertexAITokenCounter()
|
||||
model: Final = "claude-opus-5-5"
|
||||
messages: Final = [{"role": "user", "content": "Hello"}]
|
||||
system: Final = "Follow the system instructions"
|
||||
tools: Final = [
|
||||
{
|
||||
"name": "lookup",
|
||||
"description": "Look up a value",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
}
|
||||
]
|
||||
deployment: Final = {
|
||||
"litellm_params": {
|
||||
"vertex_project": "test-project",
|
||||
"vertex_location": "us-east5",
|
||||
}
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(handler, "get_async_httpx_client", return_value=fake_http_client),
|
||||
patch.object(
|
||||
handler.VertexAIPartnerModelsTokenCounter,
|
||||
"_ensure_access_token_async",
|
||||
new=AsyncMock(return_value=("fake-token", "test-project")),
|
||||
),
|
||||
):
|
||||
with_optional_fields: Final = await counter.count_tokens(
|
||||
model_to_use=model,
|
||||
messages=messages,
|
||||
contents=None,
|
||||
deployment=deployment,
|
||||
system=system,
|
||||
tools=tools,
|
||||
)
|
||||
without_optional_fields: Final = await counter.count_tokens(
|
||||
model_to_use=model,
|
||||
messages=messages,
|
||||
contents=None,
|
||||
deployment=deployment,
|
||||
)
|
||||
|
||||
assert fake_http_client.posted_bodies == (
|
||||
{"model": model, "messages": messages, "system": system, "tools": tools},
|
||||
{"model": model, "messages": messages},
|
||||
)
|
||||
assert with_optional_fields is not None
|
||||
assert with_optional_fields.total_tokens == 37
|
||||
assert without_optional_fields is not None
|
||||
assert without_optional_fields.total_tokens == 37
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("provider_failure", "expected_status", "expected_message"),
|
||||
[
|
||||
("http_400", 400, 'tools.0: Input tag "function" does not match any of the expected tags'),
|
||||
("credentials", 500, "could not resolve credentials"),
|
||||
],
|
||||
)
|
||||
async def test_vertex_ai_token_counter_returns_partner_provider_error_as_value(
|
||||
provider_failure: str, expected_status: int, expected_message: str
|
||||
):
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens import handler
|
||||
|
||||
class RejectingHttpClient:
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
json: dict[str, object],
|
||||
timeout: float,
|
||||
) -> None:
|
||||
request: Final = httpx.Request("POST", url)
|
||||
response: Final = httpx.Response(400, text=expected_message, request=request)
|
||||
raise MaskedHTTPStatusError(
|
||||
httpx.HTTPStatusError("Client error '400 Bad Request'", request=request, response=response),
|
||||
message=expected_message,
|
||||
text=expected_message,
|
||||
)
|
||||
|
||||
access_token: Final = (
|
||||
AsyncMock(side_effect=ValueError(expected_message))
|
||||
if provider_failure == "credentials"
|
||||
else AsyncMock(return_value=("fake-token", "test-project"))
|
||||
)
|
||||
with (
|
||||
patch.object(handler, "get_async_httpx_client", return_value=RejectingHttpClient()),
|
||||
patch.object(handler.VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", new=access_token),
|
||||
):
|
||||
result: Final = await VertexAITokenCounter().count_tokens(
|
||||
model_to_use="claude-opus-5-5",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"vertex_project": "test-project", "vertex_location": "us-east5"}},
|
||||
request_model="vertex-claude",
|
||||
tools=[{"type": "function", "function": {"name": "lookup", "parameters": {}}}],
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == expected_status
|
||||
assert result.error_message is not None
|
||||
assert expected_message in result.error_message
|
||||
assert result.total_tokens == 0
|
||||
assert result.request_model == "vertex-claude"
|
||||
assert result.tokenizer_type == "vertex_ai_partner_models"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_ai_token_counter_uses_count_tokens_location():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue