fix(proxy): narrow token counter exception handling

This commit is contained in:
Atharvanair09 2026-07-22 23:43:10 +05:30
parent 415b5c09fe
commit 22e00a6569
2 changed files with 17 additions and 9 deletions

View file

@ -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"):

View file

@ -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