This commit is contained in:
Atharva Nair 2026-08-27 18:36:30 -05:00 • committed by GitHub
commit 0f9a3acf4f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 213 additions and 37 deletions

View file

@ -12057,6 +12057,8 @@ async def _try_provider_token_count(
system: str | None = 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:
@ -12069,6 +12071,22 @@ async def _try_provider_token_count(
tools=tools,
system=system,
)
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"):
status_code = e.response.status_code
status_code = status_code or 500
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."
except httpx.HTTPStatusError as e:
error_message: Final = getattr(e, "message", None) or str(e)
status_code: Final = getattr(e, "status_code", None) or e.response.status_code
@ -12078,6 +12096,7 @@ async def _try_provider_token_count(
param="model",
code=status_code,
)
return None
if result is not None and result.error is True:
if litellm.disable_token_counter is True:
raise ProxyException(

View file

@ -997,11 +997,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
@ -1047,19 +1046,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
@ -1360,3 +1359,168 @@ 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
@pytest.mark.asyncio
async def test_try_provider_token_count_generic_error_fallback():
"""
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
try:
mock_provider_counter = MagicMock()
mock_provider_counter.should_use_token_counting_api.return_value = True
# 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(
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 APIError 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_generic_error_disabled():
"""
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
try:
mock_provider_counter = MagicMock()
mock_provider_counter.should_use_token_counting_api.return_value = True
# 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:
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"
)
# Status code extracted from APIError
assert exc_info.value.code == "500"
finally:
litellm.disable_token_counter = original_disable

View file

@ -67,33 +67,26 @@ class TestAgentCoreAcceptHeader:
"""
End-to-end test: verify Accept header appears in the final HTTP request
when using JWT auth through litellm.completion().
No exception swallowing: if completion() raises (for example because the
injected client was silently ignored and a real network call was made),
the test must fail with that error, not a misleading mock assertion.
"""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"result": {"role": "assistant", "content": [{"text": "agent reply"}]}
}
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
)
except Exception as e:
print(e)
raise
with patch.object(client, "post", return_value=mock_response) as mock_post:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
)
mock_post.assert_called_once()
headers = mock_post.call_args.kwargs["headers"]
assert headers["Accept"] == "application/json, text/event-stream"
assert response.choices[0].message.content == "agent reply"
mock_post.assert_called_once()
headers = mock_post.call_args.kwargs["headers"]
assert "Accept" in headers
assert headers["Accept"] == "application/json, text/event-stream"
class TestAgentCoreJsonResponseParsing: