diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py index 5d94cd42129..ddef3810bf3 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -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") diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py index 6487ea25f21..53aa07d8c5d 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py @@ -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"