fix(proxy): log mid-stream /v1/messages failures as failures with partial usage

A provider read timeout after the 200 was already committed on a streamed
/v1/messages request used to run the success logging path, so the failure
callbacks never fired and the failure metrics stayed flat. The pass-through
stream handler and the Bedrock relay iterator now dispatch the failure
handlers instead, with the usage and cost of the chunks already delivered
stashed on the logging object so the failure row still bills them.
This commit is contained in:
mateo-berri 2026-09-03 10:25:19 -07:00
parent 4990f06acc
commit cf7abf8136
7 changed files with 431 additions and 140 deletions

View file

@ -1888,6 +1888,11 @@ class Logging(LiteLLMLoggingBaseClass):
**kwargs,
)
def record_partial_usage_for_failure(self, usage: Usage, response_cost: float) -> None:
"""Stash what an interrupted stream already consumed so the failure log bills it instead of zero."""
self.model_call_details["combined_usage_object"] = usage
self.model_call_details["response_cost"] = response_cost
async def dispatch_failure_handlers(
self,
exception: Exception,

View file

@ -1,6 +1,6 @@
import asyncio
import json
from collections.abc import AsyncIterator, Mapping
from collections.abc import AsyncIterator, Mapping, Sequence
from datetime import datetime
from typing import Any, Final, Protocol, runtime_checkable
@ -176,17 +176,6 @@ def _try_claim_detached_drain_slot() -> bool:
return True
def _exception_left_unconsumed(queue: "asyncio.Queue[bytes | None | BaseException]", exc: BaseException) -> bool:
"""After client detach the relay never reads the queue again, so drain it here.
The forwarded exception still sitting in the queue means the relay tore
down before re-raising it, so the proxy's failure handling never ran and
the caller must salvage spend itself.
"""
remaining: Final = tuple(queue.get_nowait() for _ in range(queue.qsize()))
return any(item is exc for item in remaining)
def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
@ -671,27 +660,25 @@ class BaseAnthropicMessagesStreamingIterator:
self,
queue: "asyncio.Queue[bytes | None | BaseException]",
client_detached: "asyncio.Event",
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _bill_collected_chunks
exc: BaseException,
collected_chunks: Sequence[bytes],
exc: Exception,
) -> None:
"""Forward a provider error to a still-connected client, else salvage partial spend.
"""Forward a provider error to a still-connected client and log the request as failed.
Handing the original exception to the client-facing generator lets it
re-raise so the proxy's failure handling keeps the provider status and
owns logging (no success-bill). If the client already went away, or
disconnects before ever consuming the queued exception, no failure hook
runs, so bill the partial instead of dropping the request.
The relay re-raises the forwarded exception so the proxy's failure hook
keeps the provider status; the logging object's failure handlers fire
here either way, carrying the partial usage the provider already
billed, so a client that left before consuming the exception still
gets a failure row rather than a success one.
"""
from litellm._logging import verbose_proxy_logger
from litellm.proxy.pass_through_endpoints.streaming_handler import PassThroughStreamingHandler
if not client_detached.is_set() and await self._enqueue_for_client(queue, client_detached, exc):
await client_detached.wait()
if not _exception_left_unconsumed(queue, exc):
return
verbose_proxy_logger.warning(
"async_sse_wrapper upstream pump failed after client disconnect (%d chunks): %s(%s)",
len(collected_chunks),
type(exc).__name__,
exc,
if not client_detached.is_set():
await self._enqueue_for_client(queue, client_detached, exc)
PassThroughStreamingHandler.schedule_stream_failure_logging(
litellm_logging_obj=self.litellm_logging_obj,
endpoint_type=EndpointType.ANTHROPIC,
request_body=self.request_body or {},
raw_bytes=collected_chunks,
exception=exc,
)
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)

View file

@ -37,6 +37,7 @@ from litellm.types.utils import (
Message,
ModelResponse,
TextCompletionResponse,
Usage,
)
if TYPE_CHECKING:
@ -147,6 +148,134 @@ class AnthropicPassthroughLoggingHandler:
return model_group.removeprefix("passthrough/")
return model
@staticmethod
def _resolve_logged_model(
litellm_logging_obj: LiteLLMLoggingObj,
request_body: Mapping[str, object],
all_chunks: Sequence[str | bytes],
) -> str:
request_model: Final = request_body.get("model")
logged_model: Final = (
request_model
if isinstance(request_model, str) and request_model
else str(litellm_logging_obj.model_call_details.get("model") or "")
)
if logged_model and logged_model != "unknown":
return logged_model
return AnthropicPassthroughLoggingHandler._extract_model_from_anthropic_chunks(all_chunks) or logged_model
@staticmethod
def _usage_only_response_or_none(
all_chunks: Sequence[str | bytes], model: str, speed: str | None
) -> ModelResponse | None:
try:
return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=all_chunks, model=model, speed=speed
)
except Exception as e:
verbose_proxy_logger.warning("Anthropic passthrough: usage-only fallback failed (model=%s): %s", model, e)
return None
@staticmethod
def _assemble_streaming_response(
all_chunks: Sequence[str | bytes],
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
speed: str | None,
) -> ModelResponse | TextCompletionResponse | None:
try:
assembled: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
speed=speed,
)
except Exception as e:
verbose_proxy_logger.warning(
"Anthropic passthrough: stream assembly raised (model=%s): %s; falling "
"back to usage-only cost from raw SSE events.",
model,
e,
)
return AnthropicPassthroughLoggingHandler._usage_only_response_or_none(all_chunks, model, speed)
if assembled is not None:
return assembled
return AnthropicPassthroughLoggingHandler._usage_only_response_or_none(all_chunks, model, speed)
@staticmethod
def _build_streaming_response_for_logging(
litellm_logging_obj: LiteLLMLoggingObj,
request_body: Mapping[str, object],
all_chunks: Sequence[str | bytes],
model: str,
) -> ModelResponse | TextCompletionResponse | None:
response: Final = AnthropicPassthroughLoggingHandler._assemble_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
speed=AnthropicPassthroughLoggingHandler._cost_relevant_speed(request_body),
)
if response is None:
return None
AnthropicPassthroughLoggingHandler._recover_interrupted_stream_output_tokens(
response=response, all_chunks=all_chunks, model=model
)
return response
@staticmethod
def record_partial_usage_for_failure(
litellm_logging_obj: LiteLLMLoggingObj,
request_body: Mapping[str, object],
all_chunks: Sequence[str | bytes],
) -> None:
if not all_chunks:
return
model: Final = AnthropicPassthroughLoggingHandler._resolve_logged_model(
litellm_logging_obj, request_body, all_chunks
)
partial_response: Final = AnthropicPassthroughLoggingHandler._build_streaming_response_for_logging(
litellm_logging_obj=litellm_logging_obj, request_body=request_body, all_chunks=all_chunks, model=model
)
usage: Final = cast(Usage | None, getattr(partial_response, "usage", None))
if partial_response is None or usage is None:
return
try:
response_cost: Final = AnthropicPassthroughLoggingHandler._compute_response_cost(
litellm_model_response=partial_response,
model=AnthropicPassthroughLoggingHandler._resolve_costing_model(model, litellm_logging_obj),
logging_obj=litellm_logging_obj,
)
except Exception as e:
verbose_proxy_logger.warning(
"Anthropic passthrough: could not cost the partial usage of a failed stream (model=%s): %s", model, e
)
return
litellm_logging_obj.record_partial_usage_for_failure(usage=usage, response_cost=response_cost)
@staticmethod
def _compute_response_cost(
litellm_model_response: ModelResponse | TextCompletionResponse,
model: str,
logging_obj: LiteLLMLoggingObj,
) -> float:
if logging_obj.model_call_details.get("cache_hit") is True:
return 0.0
custom_llm_provider: Final = logging_obj.model_call_details.get("custom_llm_provider")
model_for_cost: Final = (
f"{custom_llm_provider}/{model}"
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/")
else model
)
return litellm.completion_cost(
completion_response=litellm_model_response,
model=model_for_cost,
custom_llm_provider=custom_llm_provider,
custom_pricing=use_custom_pricing_for_model(
litellm_params=(logging_obj.litellm_params if hasattr(logging_obj, "litellm_params") else None)
),
router_model_id=logging_obj.get_router_model_id(),
)
@staticmethod
def _extract_model_from_anthropic_chunks(
all_chunks: Sequence[str | bytes],
@ -263,31 +392,9 @@ class AnthropicPassthroughLoggingHandler:
if logging_obj.model_call_details.get("stream") is True:
logging_obj.model_call_details["complete_streaming_response"] = litellm_model_response
try:
# Get custom_llm_provider from logging object if available (e.g., azure_ai for Azure Anthropic)
custom_llm_provider: Final = logging_obj.model_call_details.get("custom_llm_provider")
model = AnthropicPassthroughLoggingHandler._resolve_costing_model(model, logging_obj)
# Prepend custom_llm_provider to model if not already present
model_for_cost = model
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
model_for_cost = f"{custom_llm_provider}/{model}"
router_model_id: Final = logging_obj.get_router_model_id()
custom_pricing: Final = use_custom_pricing_for_model(
litellm_params=(logging_obj.litellm_params if hasattr(logging_obj, "litellm_params") else None)
)
response_cost: Final = (
0.0
if logging_obj.model_call_details.get("cache_hit") is True
else litellm.completion_cost(
completion_response=litellm_model_response,
model=model_for_cost,
custom_llm_provider=custom_llm_provider,
custom_pricing=custom_pricing,
router_model_id=router_model_id,
)
response_cost: Final = AnthropicPassthroughLoggingHandler._compute_response_cost(
litellm_model_response=litellm_model_response, model=model, logging_obj=logging_obj
)
kwargs["response_cost"] = response_cost
@ -342,57 +449,12 @@ class AnthropicPassthroughLoggingHandler:
- Logs in litellm callbacks
"""
speed: Final = AnthropicPassthroughLoggingHandler._cost_relevant_speed(request_body)
model = request_body.get("model", "")
# Check if it's available in the logging object
if (
not model
and hasattr(litellm_logging_obj, "model_call_details")
and litellm_logging_obj.model_call_details.get("model")
):
model = cast(str, litellm_logging_obj.model_call_details.get("model"))
if not model or model == "unknown":
chunk_model: Final = AnthropicPassthroughLoggingHandler._extract_model_from_anthropic_chunks(all_chunks)
if chunk_model:
model = chunk_model
try:
complete_streaming_response = AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
speed=speed,
)
except Exception as e:
# stream_chunk_builder re-raises assembly failures (as litellm.APIError)
# on large agentic tool-use / thinking streams; treat that the same as a
# None result so the usage-only fallback below still recovers cost
verbose_proxy_logger.warning(
"Anthropic passthrough: stream assembly raised (model=%s): %s; falling "
"back to usage-only cost from raw SSE events.",
model,
e,
)
complete_streaming_response = None
if complete_streaming_response is None:
# stream_chunk_builder cannot always reassemble large agentic streams, but
# Anthropic still emits token usage in the message_start / message_delta SSE
# events regardless of content shape; recover usage-only so cost is tracked.
# Guard it too: a raise here would defeat the point and drop the request
try:
complete_streaming_response = AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=all_chunks,
model=model,
speed=speed,
)
except Exception as e:
verbose_proxy_logger.warning(
"Anthropic passthrough: usage-only fallback failed (model=%s): %s",
model,
e,
)
complete_streaming_response = None
model: Final = AnthropicPassthroughLoggingHandler._resolve_logged_model(
litellm_logging_obj, request_body, all_chunks
)
complete_streaming_response: Final = AnthropicPassthroughLoggingHandler._build_streaming_response_for_logging(
litellm_logging_obj=litellm_logging_obj, request_body=request_body, all_chunks=all_chunks, model=model
)
if complete_streaming_response is None:
verbose_proxy_logger.error(
"Unable to build complete streaming response for Anthropic passthrough endpoint, not logging..."
@ -401,11 +463,6 @@ class AnthropicPassthroughLoggingHandler:
"result": None,
"kwargs": {},
}
AnthropicPassthroughLoggingHandler._recover_interrupted_stream_output_tokens(
response=complete_streaming_response,
all_chunks=all_chunks,
model=model,
)
kwargs: Final = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
litellm_model_response=complete_streaming_response,
model=model,

View file

@ -1,4 +1,5 @@
from collections.abc import Coroutine
import traceback
from collections.abc import Coroutine, Mapping, Sequence
from datetime import datetime
from typing import Final, Protocol
@ -50,6 +51,27 @@ class PassThroughStreamingHandler:
if litellm_logging_obj.completion_start_time is None:
litellm_logging_obj._update_completion_start_time(completion_start_time=datetime.now())
@staticmethod
def schedule_stream_failure_logging(
litellm_logging_obj: LiteLLMLoggingObj,
endpoint_type: EndpointType,
request_body: Mapping[str, object],
raw_bytes: Sequence[bytes],
exception: Exception,
) -> None:
if endpoint_type == EndpointType.ANTHROPIC:
AnthropicPassthroughLoggingHandler.record_partial_usage_for_failure(
litellm_logging_obj=litellm_logging_obj, request_body=request_body, all_chunks=raw_bytes
)
try:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=litellm_logging_obj.dispatch_failure_handlers(
exception, traceback.format_exc(), prefer_async_handlers=True
)
)
except Exception as e:
verbose_proxy_logger.error("Error scheduling stream failure logging: %s", e)
@staticmethod
async def chunk_processor(
response: httpx.Response,
@ -132,9 +154,9 @@ class PassThroughStreamingHandler:
# coroutine on logging_obj instead of enqueueing now, so
# ProxyLogging._fire_deferred_stream_logging fires it after
# guardrail end-of-stream blocks populate guardrail_information.
# Disconnect/exception paths skip this and fall through to the
# immediate enqueue in ``finally`` to keep partial billing
# (LIT-2642).
# Disconnect paths skip this and fall through to the immediate
# enqueue in ``finally`` to keep partial billing (LIT-2642);
# upstream exceptions log a failure instead (LIT-3798).
if (
getattr(litellm_logging_obj, "_on_deferred_stream_complete", None) is not None
and raw_bytes
@ -144,6 +166,15 @@ class PassThroughStreamingHandler:
litellm_logging_obj._deferred_stream_complete_args = (_build_logging_coroutine(),)
except Exception as e:
verbose_proxy_logger.error("Error in chunk_processor: %s", e)
if response.status_code < 400:
logging_scheduled = True
PassThroughStreamingHandler.schedule_stream_failure_logging(
litellm_logging_obj=litellm_logging_obj,
endpoint_type=endpoint_type,
request_body=request_body or {},
raw_bytes=raw_bytes,
exception=e,
)
raise
finally:
# GeneratorExit (raised on client disconnect) is not caught by

View file

@ -5,6 +5,7 @@ from datetime import datetime
import pytest
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.experimental_pass_through.messages import streaming_iterator as streaming_iterator_module
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
@ -32,7 +33,16 @@ class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator):
self.logging_call_count += 1
def _make_logging_obj(test_name: str) -> LiteLLMLoggingObj:
class _FailureRecorder(CustomLogger):
def __init__(self):
super().__init__()
self.failure_kwargs: list = []
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
self.failure_kwargs.append(kwargs)
def _make_logging_obj(test_name: str, failure_recorder: _FailureRecorder | None = None) -> LiteLLMLoggingObj:
return LiteLLMLoggingObj(
model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
@ -41,9 +51,19 @@ def _make_logging_obj(test_name: str) -> LiteLLMLoggingObj:
start_time=datetime.now(),
litellm_call_id=test_name,
function_id=test_name,
dynamic_async_failure_callbacks=[failure_recorder] if failure_recorder is not None else None,
)
async def _wait_for_failure_event(recorder: _FailureRecorder) -> dict:
for _ in range(300):
if recorder.failure_kwargs:
break
await asyncio.sleep(0.01)
assert len(recorder.failure_kwargs) == 1, "expected exactly one failure event"
return recorder.failure_kwargs[0]
def _make_iterator(test_name: str) -> BaseAnthropicMessagesStreamingIterator:
return BaseAnthropicMessagesStreamingIterator(
litellm_logging_obj=_make_logging_obj(test_name),
@ -539,7 +559,8 @@ async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client():
before message_stop must propagate the ORIGINAL provider exception to a
still-connected client, so the proxy's failure handling keeps the
provider-specific status. The pump must not swallow it into a generic
api_error event + normal termination.
api_error event + normal termination, and the request is logged as a
failure carrying the partial usage, never as a success.
"""
async def _failing_stream():
@ -547,8 +568,9 @@ async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client():
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}
raise _ProviderStreamError("bedrock stream blew up", status_code=529)
recorder = _FailureRecorder()
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_reraises_upstream_error"),
litellm_logging_obj=_make_logging_obj("test_reraises_upstream_error", recorder),
request_body={},
)
@ -561,18 +583,23 @@ async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client():
with pytest.raises(_ProviderStreamError) as excinfo:
await _drain()
failure_kwargs = await _wait_for_failure_event(recorder)
assert excinfo.value.status_code == 529
assert received
assert not any(c.startswith(b"event: error\n") for c in received)
assert iterator.logged_chunks == []
assert failure_kwargs["standard_logging_object"]["status"] == "failure"
assert failure_kwargs["standard_logging_object"]["prompt_tokens"] == 52
@pytest.mark.asyncio
async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_disconnect():
async def test_async_sse_wrapper_logs_failure_on_upstream_error_after_disconnect():
"""
When the upstream errors AFTER the client has already disconnected there is
no live client to re-raise to and no failure hook will run, so the pump
salvages partial spend from what it collected instead of dropping the row.
no live client to re-raise to and no proxy failure hook will run, so the
pump logs the failure itself with the partial usage it collected; it must
never bill the broken stream as a success.
"""
tail_gated = asyncio.Event()
@ -582,8 +609,9 @@ async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_
await tail_gated.wait()
raise _ProviderStreamError("late failure", status_code=500)
recorder = _FailureRecorder()
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_salvage_partial_on_late_error"),
litellm_logging_obj=_make_logging_obj("test_failure_logged_on_late_error", recorder),
request_body={},
)
@ -592,24 +620,23 @@ async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_
await gen.aclose() # client disconnects before the upstream error
tail_gated.set() # let the upstream raise now, after disconnect
for _ in range(100):
if iterator.logged_chunks:
break
await asyncio.sleep(0.01)
failure_kwargs = await _wait_for_failure_event(recorder)
assert len(received) == 2
assert iterator.logged_chunks == received
assert iterator.logging_call_count == 0
assert failure_kwargs["standard_logging_object"]["status"] == "failure"
assert failure_kwargs["standard_logging_object"]["prompt_tokens"] == 52
assert isinstance(failure_kwargs["exception"], _ProviderStreamError)
@pytest.mark.asyncio
async def test_async_sse_wrapper_salvages_spend_when_queued_error_is_never_consumed():
async def test_async_sse_wrapper_logs_failure_when_queued_error_is_never_consumed():
"""
When the upstream errors while the client is still connected, the pump
forwards the exception through the queue expecting the relay to re-raise it
into the proxy's failure handling. If the client disconnects before
consuming that queued exception, the handoff never happens and no failure
hook runs, so the pump must notice the unconsumed exception at teardown and
salvage partial spend instead of dropping the row entirely.
forwards the exception through the queue for the relay to re-raise. If the
client disconnects before consuming that queued exception, no proxy failure
hook runs, so the failure logged by the pump itself is the only record of
the request; it must be a failure row, not a salvaged success.
"""
upstream_errored = asyncio.Event()
@ -619,8 +646,9 @@ async def test_async_sse_wrapper_salvages_spend_when_queued_error_is_never_consu
upstream_errored.set()
raise _ProviderStreamError("mid-stream failure", status_code=500)
recorder = _FailureRecorder()
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_salvage_on_unconsumed_queued_error"),
litellm_logging_obj=_make_logging_obj("test_failure_logged_on_unconsumed_queued_error", recorder),
request_body={},
)
@ -629,13 +657,12 @@ async def test_async_sse_wrapper_salvages_spend_when_queued_error_is_never_consu
await upstream_errored.wait() # exception is now queued behind the consumed chunks
await gen.aclose() # client disconnects without ever consuming the queued exception
for _ in range(100):
if iterator.logged_chunks:
break
await asyncio.sleep(0.01)
failure_kwargs = await _wait_for_failure_event(recorder)
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == received
assert len(received) == 2
assert iterator.logging_call_count == 0
assert failure_kwargs["standard_logging_object"]["status"] == "failure"
assert failure_kwargs["standard_logging_object"]["prompt_tokens"] == 52
@pytest.mark.asyncio

View file

@ -2441,3 +2441,78 @@ class TestAnthropicPassthroughFastMode:
assert served_standard.usage.speed == "standard"
assert self._cost(served_standard) == pytest.approx(self._cost(standard))
class TestRecordPartialUsageForFailure:
"""A stream that dies mid-way still carries the usage the provider billed in
message_start; the failure row must keep it and its cost instead of logging
a zero-cost failure (or, worse, a success)."""
@staticmethod
def _sse(event, data):
return f"event: {event}\ndata: {json.dumps(data)}\n\n".encode()
@staticmethod
def _make_logging_obj() -> LiteLLMLoggingObj:
return LiteLLMLoggingObj(
model="claude-sonnet-5",
messages=[{"role": "user", "content": "hello"}],
stream=True,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id="test-partial-usage-failure",
function_id="test-partial-usage-failure",
)
def _interrupted_chunks(self):
return [
self._sse(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_abc",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 52, "output_tokens": 1},
},
},
),
self._sse(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
self._sse(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}},
),
]
def test_stashes_partial_usage_and_cost_from_interrupted_stream(self):
logging_obj = self._make_logging_obj()
AnthropicPassthroughLoggingHandler.record_partial_usage_for_failure(
litellm_logging_obj=logging_obj,
request_body={"model": "claude-sonnet-5", "stream": True},
all_chunks=self._interrupted_chunks(),
)
usage = logging_obj.model_call_details["combined_usage_object"]
assert usage.prompt_tokens == 52
assert logging_obj.model_call_details["response_cost"] > 0
def test_leaves_logging_obj_untouched_when_nothing_streamed(self):
logging_obj = self._make_logging_obj()
AnthropicPassthroughLoggingHandler.record_partial_usage_for_failure(
litellm_logging_obj=logging_obj,
request_body={"model": "claude-sonnet-5", "stream": True},
all_chunks=[],
)
assert "combined_usage_object" not in logging_obj.model_call_details
assert "response_cost" not in logging_obj.model_call_details

View file

@ -9,6 +9,8 @@ import httpx
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy.pass_through_endpoints.streaming_handler import (
PassThroughStreamingHandler,
@ -632,3 +634,110 @@ async def test_chunk_processor_enqueues_immediately_on_disconnect_even_when_arme
mock_enqueue.assert_called_once()
assert logging_obj._deferred_stream_complete_args is None
class _EventRecorder(CustomLogger):
def __init__(self):
super().__init__()
self.failure_kwargs = []
self.success_kwargs = []
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
self.failure_kwargs.append(kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.success_kwargs.append(kwargs)
def _anthropic_sse(event: str, payload: dict) -> bytes:
return f"event: {event}\ndata: {json.dumps(payload)}\n\n".encode()
def _anthropic_stream_that_times_out_mid_stream():
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
async def _aiter_bytes():
yield _anthropic_sse(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 52, "output_tokens": 1},
},
},
)
yield _anthropic_sse(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
)
yield _anthropic_sse(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}},
)
raise httpx.ReadTimeout("Timeout on reading data from socket")
mock.aiter_bytes = _aiter_bytes
return mock
@pytest.mark.asyncio
async def test_chunk_processor_logs_failure_not_success_on_mid_stream_exception():
"""A stream that dies after the first chunks is a failed request: the failure
callbacks must fire once with the partial usage and cost, and the success
routing must never run for it."""
recorder = _EventRecorder()
logging_obj = LiteLLMLoggingObj(
model="claude-sonnet-5",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id="test-mid-stream-timeout",
function_id="test-mid-stream-timeout",
dynamic_async_success_callbacks=[recorder],
dynamic_async_failure_callbacks=[recorder],
)
success_routes = []
async def _record_success_route(**kwargs):
success_routes.append(kwargs)
received = []
async def _consume_stream():
async for chunk in PassThroughStreamingHandler.chunk_processor(
response=_anthropic_stream_that_times_out_mid_stream(),
request_body={"model": "claude-sonnet-5", "stream": True},
litellm_logging_obj=logging_obj,
endpoint_type=EndpointType.ANTHROPIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
url_route="/v1/messages",
route_streaming_logging=_record_success_route,
):
received.append(chunk)
with pytest.raises(httpx.ReadTimeout):
await _consume_stream()
for _ in range(300):
if recorder.failure_kwargs:
break
await asyncio.sleep(0.01)
assert len(received) == 3
assert success_routes == []
assert recorder.success_kwargs == []
assert len(recorder.failure_kwargs) == 1
failure_payload = recorder.failure_kwargs[0]["standard_logging_object"]
assert failure_payload["status"] == "failure"
assert failure_payload["prompt_tokens"] == 52
assert failure_payload["response_cost"] > 0
assert isinstance(recorder.failure_kwargs[0]["exception"], httpx.ReadTimeout)