From dd70a19dbc4084568cfee15ea847b39259760cb7 Mon Sep 17 00:00:00 2001 From: Atharvanair09 Date: Wed, 22 Jul 2026 21:12:39 +0530 Subject: [PATCH] fix(proxy): fall back to local tokenizer on provider token count errors --- litellm/proxy/proxy_server.py | 16 ++- .../test_proxy_token_counter.py | 112 +++++++++++++++--- 2 files changed, 108 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 32845763f22..0d425fac1a0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10784,12 +10784,18 @@ async def _try_provider_token_count( except httpx.HTTPStatusError as e: error_message = getattr(e, "message", None) or str(e) status_code = getattr(e, "status_code", None) or e.response.status_code - raise ProxyException( - message=error_message, - type="token_counting_error", - param="model", - code=status_code, + if litellm.disable_token_counter is True: + raise ProxyException( + message=error_message, + type="token_counting_error", + param="model", + code=status_code, + ) + verbose_proxy_logger.warning( + f"Provider token counting failed ({status_code}): {error_message}. " + "Falling back to local tokenizer." ) + return None if result is not None and result.error is True: if litellm.disable_token_counter is True: raise ProxyException( diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index 1079a5228a1..e7a25af2659 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -1001,11 +1001,10 @@ async def test_bedrock_handler_httpx_error_status_code_propagation(): @pytest.mark.asyncio -async def test_token_counter_httpx_status_error_raises_proxy_exception(): +async def test_token_counter_httpx_status_error_falls_back(): """ When provider_counter.count_tokens() raises httpx.HTTPStatusError, - the token_counter endpoint should catch it and raise a ProxyException - with the upstream status code and error message. + the token_counter endpoint should catch it and fall back to the local tokenizer. """ upstream_status = 429 @@ -1051,19 +1050,19 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception(): ) litellm.proxy.proxy_server.llm_router = mock_router - with pytest.raises(ProxyException) as exc_info: - await token_counter( - request=TokenCountRequest( - model="claude-4-6-sonnet", - messages=[{"role": "user", "content": "hello"}], - ), - call_endpoint=True, - ) + # When disabled=False (default), it should fall back to local token counting + response_obj = await token_counter( + request=TokenCountRequest( + model="claude-4-6-sonnet", + messages=[{"role": "user", "content": "hello"}], + ), + call_endpoint=True, + ) - assert exc_info.value.code == str(upstream_status) - assert upstream_message in exc_info.value.message - assert exc_info.value.type == "token_counting_error" - assert exc_info.value.param == "model" + assert response_obj.total_tokens > 0 + assert response_obj.request_model == "claude-4-6-sonnet" + assert response_obj.model_used == "claude-4-6-sonnet" + assert response_obj.tokenizer_type != "vertex_ai" finally: litellm.proxy.proxy_server._get_provider_token_counter = ( original_get_provider_token_counter @@ -1364,3 +1363,86 @@ async def test_anthropic_endpoint_429_rate_limit_error_format(): finally: anthropic_endpoints._read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter + + +@pytest.mark.asyncio +async def test_try_provider_token_count_http_error_fallback(): + """ + Test that _try_provider_token_count catches httpx.HTTPStatusError + and returns None (falling back to local tokenizer) when litellm.disable_token_counter != True. + """ + from litellm.proxy.proxy_server import _try_provider_token_count + from unittest.mock import AsyncMock, MagicMock + import litellm + import httpx + + original_disable = getattr(litellm, "disable_token_counter", False) + litellm.disable_token_counter = False + + try: + mock_provider_counter = MagicMock() + mock_provider_counter.should_use_token_counting_api.return_value = True + + # Make count_tokens raise HTTPStatusError + request = httpx.Request("POST", "https://example.com") + response = httpx.Response(500, request=request) + error = httpx.HTTPStatusError("Internal Server Error", request=request, response=response) + + mock_provider_counter.count_tokens = AsyncMock(side_effect=error) + + result = await _try_provider_token_count( + provider_counter=mock_provider_counter, + custom_llm_provider="openai", + model_to_use="gpt-4", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment=None, + request_model="gpt-4" + ) + + # It should return None on HTTPStatusError instead of raising Exception + assert result is None + finally: + litellm.disable_token_counter = original_disable + + +@pytest.mark.asyncio +async def test_try_provider_token_count_http_error_disabled(): + """ + Test that _try_provider_token_count raises ProxyException on httpx.HTTPStatusError + when litellm.disable_token_counter == True. + """ + from litellm.proxy.proxy_server import _try_provider_token_count + from litellm.proxy._types import ProxyException + from unittest.mock import AsyncMock, MagicMock + import litellm + import httpx + + original_disable = getattr(litellm, "disable_token_counter", False) + litellm.disable_token_counter = True + + try: + mock_provider_counter = MagicMock() + mock_provider_counter.should_use_token_counting_api.return_value = True + + # Make count_tokens raise HTTPStatusError + request = httpx.Request("POST", "https://example.com") + response = httpx.Response(500, request=request) + error = httpx.HTTPStatusError("Internal Server Error", request=request, response=response) + + mock_provider_counter.count_tokens = AsyncMock(side_effect=error) + + with pytest.raises(ProxyException) as exc_info: + await _try_provider_token_count( + provider_counter=mock_provider_counter, + custom_llm_provider="openai", + model_to_use="gpt-4", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment=None, + request_model="gpt-4" + ) + + assert exc_info.value.code == "500" + finally: + litellm.disable_token_counter = original_disable