diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3856a1f10b8..162482b366f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10769,6 +10769,8 @@ async def _try_provider_token_count( system: Optional[str] = None, ) -> Optional["TokenCountResponse"]: """Attempt provider-specific token counting. Returns result on success, None to fall through to local counting.""" + import openai + if not provider_counter.should_use_token_counting_api(custom_llm_provider=custom_llm_provider): return None try: @@ -10781,7 +10783,7 @@ async def _try_provider_token_count( tools=tools, system=system, ) - except Exception as e: + except (httpx.HTTPError, openai.OpenAIError) as e: error_message = getattr(e, "message", None) or str(e) status_code = getattr(e, "status_code", None) if status_code is None and hasattr(e, "response") and hasattr(e.response, "status_code"): diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index 232aae0cb10..b6980f189a8 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -1451,12 +1451,14 @@ async def test_try_provider_token_count_http_error_disabled(): @pytest.mark.asyncio async def test_try_provider_token_count_generic_error_fallback(): """ - Test that _try_provider_token_count catches a generic Exception + Test that _try_provider_token_count catches a provider API exception (like APIError) 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 + from litellm.exceptions import APIError + import httpx original_disable = getattr(litellm, "disable_token_counter", False) litellm.disable_token_counter = False @@ -1465,8 +1467,9 @@ async def test_try_provider_token_count_generic_error_fallback(): mock_provider_counter = MagicMock() mock_provider_counter.should_use_token_counting_api.return_value = True - # Make count_tokens raise a generic Exception - error = ValueError("Some weird provider error") + # Make count_tokens raise a non-HTTP provider exception + request = httpx.Request(method="POST", url="https://api.openai.com/v1") + error = APIError(status_code=500, message="Provider internal error", llm_provider="openai", model="gpt-4", request=request) mock_provider_counter.count_tokens = AsyncMock(side_effect=error) result = await _try_provider_token_count( @@ -1479,7 +1482,7 @@ async def test_try_provider_token_count_generic_error_fallback(): request_model="gpt-4" ) - # It should return None on generic Exception instead of raising Exception + # It should return None on APIError instead of raising Exception assert result is None finally: litellm.disable_token_counter = original_disable @@ -1488,13 +1491,15 @@ async def test_try_provider_token_count_generic_error_fallback(): @pytest.mark.asyncio async def test_try_provider_token_count_generic_error_disabled(): """ - Test that _try_provider_token_count raises ProxyException on generic Exception + Test that _try_provider_token_count raises ProxyException on provider API exception 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 + from litellm.exceptions import APIError + import httpx original_disable = getattr(litellm, "disable_token_counter", False) litellm.disable_token_counter = True @@ -1503,8 +1508,9 @@ async def test_try_provider_token_count_generic_error_disabled(): mock_provider_counter = MagicMock() mock_provider_counter.should_use_token_counting_api.return_value = True - # Make count_tokens raise a generic Exception - error = ValueError("Some weird provider error") + # Make count_tokens raise a non-HTTP provider exception + request = httpx.Request(method="POST", url="https://api.openai.com/v1") + error = APIError(status_code=500, message="Provider internal error", llm_provider="openai", model="gpt-4", request=request) mock_provider_counter.count_tokens = AsyncMock(side_effect=error) with pytest.raises(ProxyException) as exc_info: @@ -1518,7 +1524,7 @@ async def test_try_provider_token_count_generic_error_disabled(): request_model="gpt-4" ) - # Default status code for generic exceptions is 500 + # Status code extracted from APIError assert exc_info.value.code == "500" finally: litellm.disable_token_counter = original_disable