mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): narrow token counter exception handling
This commit is contained in:
parent
415b5c09fe
commit
22e00a6569
2 changed files with 17 additions and 9 deletions
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue