mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
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:
parent
4990f06acc
commit
cf7abf8136
7 changed files with 431 additions and 140 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue