fix(vertex): strip version suffix from model name in count_tokens requests (#25800)

The Vertex AI count-tokens endpoint rejects model names that include
version suffixes (@default, @20251001, etc.) with:
"claude-sonnet-4-6@default is not supported for token counting"

The same model without the suffix ("claude-sonnet-4-6") works correctly.

Strip @suffix from both the model parameter and request_data["model"]
in handle_count_tokens_request before sending to the API.

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Darien Kindlund 2026-04-15 22:24:07 -04:00 committed by Sameer Kankute
parent ed0138b50e
commit 6b2973b29a
No known key found for this signature in database
2 changed files with 88 additions and 0 deletions

View file

@ -78,6 +78,19 @@ class VertexAIPartnerModelsTokenCounter(VertexBase):
return endpoint
@staticmethod
def _strip_version_suffix(model: str) -> str:
"""
Strip version suffixes (e.g. @default, @20251001) from model names.
The Vertex AI count-tokens endpoint rejects model names that include
version suffixes for example, "claude-sonnet-4-6@default" returns
"not supported for token counting" while "claude-sonnet-4-6" works.
"""
if "@" in model:
return model.split("@")[0]
return model
async def handle_count_tokens_request(
self,
model: str,
@ -98,6 +111,15 @@ class VertexAIPartnerModelsTokenCounter(VertexBase):
Raises:
ValueError: If required parameters are missing or invalid
"""
# Strip version suffixes (@default, @20251001, etc.) — the Vertex AI
# count-tokens endpoint does not accept versioned model names.
model = self._strip_version_suffix(model)
if "model" in request_data:
request_data = {
**request_data,
"model": self._strip_version_suffix(request_data["model"]),
}
# Validate request
if "messages" not in request_data:
raise ValueError("messages required for token counting")

View file

@ -162,3 +162,69 @@ class TestCountTokensLocationResolution:
)
assert captured["vertex_location"] == "asia-southeast1"
class TestCountTokensVersionSuffixStripping:
"""Verify that version suffixes (@default, @20251001, etc.) are stripped
from model names before sending to the Vertex AI count-tokens endpoint.
The Vertex AI count-tokens API rejects versioned model names with:
"claude-sonnet-4-6@default is not supported for token counting"
while "claude-sonnet-4-6" (without suffix) works correctly.
"""
def test_strip_version_suffix_at_default(self):
counter = VertexAIPartnerModelsTokenCounter()
assert counter._strip_version_suffix("claude-sonnet-4-6@default") == "claude-sonnet-4-6"
def test_strip_version_suffix_at_date(self):
counter = VertexAIPartnerModelsTokenCounter()
assert counter._strip_version_suffix("claude-haiku-4-5@20251001") == "claude-haiku-4-5"
def test_strip_version_suffix_no_suffix(self):
counter = VertexAIPartnerModelsTokenCounter()
assert counter._strip_version_suffix("claude-sonnet-4-6") == "claude-sonnet-4-6"
@pytest.mark.asyncio
async def test_handle_count_tokens_strips_version_from_request_data(self, monkeypatch):
"""The model name in request_data sent to the API must have @suffix stripped."""
counter = VertexAIPartnerModelsTokenCounter()
captured_json = {}
async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider):
return "fake-token", "fake-project"
def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None):
return "https://fake-endpoint"
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token
)
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint
)
class FakeResponse:
status_code = 200
def json(self):
return {"input_tokens": 10}
class FakeClient:
async def post(self, url, headers=None, json=None, **kwargs):
captured_json.update(json or {})
return FakeResponse()
import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod
monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient())
await counter.handle_count_tokens_request(
model="claude-sonnet-4-6@default",
request_data={
"model": "claude-sonnet-4-6@default",
"messages": [{"role": "user", "content": "hi"}],
},
litellm_params={"vertex_location": "us-east5"},
)
# The model name sent to the API must NOT have the @default suffix
assert captured_json["model"] == "claude-sonnet-4-6"