test(proxy): move in-band error tests into the existing streaming handler test file

This commit is contained in:
Deepanshu Pal 2026-10-02 06:43:04 +05:30 • committed by GitHub
parent 28197fe0cb
commit 869a5df8cc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -1,7 +1,7 @@
import json
from collections.abc import Iterator
from datetime import datetime
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -230,3 +230,71 @@ async def test_failed_anthropic_stream_records_partial_usage_off_the_event_loop(
partial_usage = logging_obj.record_partial_usage_for_failure.call_args.kwargs["usage"]
assert partial_usage.completion_tokens > 100_000
assert_loop_stayed_free(took, lags)
CLAUDE_MODEL = "claude-fable-5"
def _anthropic_error_frame(error_type: str, message: str) -> bytes:
payload = {"type": "error", "error": {"type": error_type, "message": message}}
return f"event: error\ndata: {json.dumps(payload)}\n\n".encode()
def _anthropic_error_logging_obj() -> MagicMock:
logging_obj = MagicMock(spec=LiteLLMLoggingObj)
logging_obj.model_call_details = {"model": CLAUDE_MODEL, "stream": True}
logging_obj.optional_params = {}
logging_obj.litellm_params = {}
logging_obj.litellm_call_id = "test-call-id"
logging_obj.get_router_model_id.return_value = None
logging_obj.dispatch_success_handlers = AsyncMock()
logging_obj.dispatch_failure_handlers = AsyncMock()
return logging_obj
async def _route_anthropic_stream(logging_obj: MagicMock, raw_bytes: list[bytes]) -> None:
await PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=PassThroughEndpointLogging(),
url_route="/anthropic/v1/messages",
request_body={"model": CLAUDE_MODEL, "stream": True},
endpoint_type=EndpointType.ANTHROPIC,
start_time=datetime.now(),
raw_bytes=raw_bytes,
end_time=datetime.now(),
model=CLAUDE_MODEL,
)
await GLOBAL_LOGGING_WORKER.flush()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error_type, expected_status",
[("api_error", 500), ("overloaded_error", 503), ("rate_limit_error", 429)],
)
async def test_in_band_error_frame_is_logged_as_failure_not_success(error_type: str, expected_status: int):
logging_obj = _anthropic_error_logging_obj()
await _route_anthropic_stream(
logging_obj,
[*_interrupted_anthropic_stream(CLAUDE_MODEL, "partial answer"), _anthropic_error_frame(error_type, "boom")],
)
logging_obj.dispatch_success_handlers.assert_not_awaited()
logging_obj.dispatch_failure_handlers.assert_awaited_once()
exception = logging_obj.dispatch_failure_handlers.await_args.args[0]
assert exception.status_code == expected_status
assert "boom" in str(exception)
partial_usage = logging_obj.record_partial_usage_for_failure.call_args.kwargs["usage"]
assert partial_usage.prompt_tokens == 29
assert partial_usage.completion_tokens > 0
@pytest.mark.asyncio
async def test_stream_without_error_frame_still_logs_success():
logging_obj = _anthropic_error_logging_obj()
await _route_anthropic_stream(logging_obj, _interrupted_anthropic_stream(CLAUDE_MODEL, "partial answer"))
logging_obj.dispatch_success_handlers.assert_awaited_once()
logging_obj.dispatch_failure_handlers.assert_not_awaited()