mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 7de530c57e into 49affa7c01
This commit is contained in:
commit
0f9a3acf4f
3 changed files with 213 additions and 37 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue