mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
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:
parent
ed0138b50e
commit
6b2973b29a
2 changed files with 88 additions and 0 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue