test(service-tier): bill disconnects through the router's anthropic stream wrapper

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-25 01:28:15 +00:00
parent bf86914da4
commit 047b506c4e
3 changed files with 111 additions and 11 deletions

View file

@ -273,10 +273,9 @@ class TestServiceTierPricing:
served_tier = completed.response.service_tier
assert served_tier, f"response.completed carried no service_tier: {completed.response}"
assert served_tier in TIER_INPUT_RATES, f"no custom rate registered for served tier {served_tier!r}"
assert completed.response.id, f"response.completed carried no id: {completed.response}"
row = poll_cost_row(client.proxy, completed.response.id)
assert row is not None, f"no spend row with a cost breakdown landed for {completed.response.id}"
row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.spend is not None and r.spend > 0)
assert row is not None, f"no spend row with a cost breakdown landed for the streamed responses call on {model}"
assert row.breakdown.service_tier == served_tier, (
f"response.completed served tier {served_tier!r} but the bill records "
f"pricing basis {row.breakdown.service_tier!r}"

View file

@ -0,0 +1,87 @@
"""
Tests for AnthropicSSEStream, the object translate_completion_output_params_streaming
hands to the proxy for /v1/messages streaming. It must emit the same SSE bytes as
the wrapper's async_anthropic_sse_wrapper, propagate aclose into it, and expose the
wrapper's chunks/messages/model so disconnect-time partial billing can read them.
"""
from typing import Final
from unittest.mock import MagicMock
import pytest
from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import (
AnthropicSSEStream,
AnthropicStreamWrapper,
)
from litellm.types.utils import Delta, StreamingChoices
def _make_chunk(delta: Delta, finish_reason: str | None = None) -> MagicMock:
chunk = MagicMock()
chunk.choices = [StreamingChoices(finish_reason=finish_reason, index=0, delta=delta, logprobs=None)]
chunk.usage = None
chunk._hidden_params = {}
return chunk
class _AsyncStream:
def __init__(self, items: list[MagicMock]):
self._it = iter(items)
self.chunks = list(items)
self.messages: list[dict] = [{"role": "user", "content": "hi"}]
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self._it)
except StopIteration:
raise StopAsyncIteration
def _streamed_events() -> AnthropicSSEStream:
upstream: Final = _AsyncStream(
[
_make_chunk(Delta(content="Once")),
_make_chunk(Delta(content=" upon"), finish_reason="stop"),
]
)
wrapper: Final = AnthropicStreamWrapper(completion_stream=upstream, model="gpt-4o-mini")
wrapper._message_id = "msg_test"
return AnthropicSSEStream(wrapper)
@pytest.mark.asyncio
async def test_sse_stream_yields_identical_bytes_to_the_wrappers_sse_wrapper():
upstream_a: Final = _AsyncStream(
[_make_chunk(Delta(content="Once")), _make_chunk(Delta(content=" upon"), finish_reason="stop")]
)
wrapper_a: Final = AnthropicStreamWrapper(completion_stream=upstream_a, model="gpt-4o-mini")
wrapper_a._message_id = "msg_test"
expected: Final = [event async for event in wrapper_a.async_anthropic_sse_wrapper()]
actual: Final = [event async for event in _streamed_events()]
assert actual == expected
@pytest.mark.asyncio
async def test_sse_stream_aclose_ends_the_wrapped_stream():
stream: Final = _streamed_events()
first: Final = await stream.__anext__()
assert first.startswith(b"event: message_start")
await stream.aclose()
with pytest.raises(StopAsyncIteration):
await stream.__anext__()
def test_sse_stream_exposes_chunks_messages_and_model():
stream: Final = _streamed_events()
assert stream.model == "gpt-4o-mini"
assert stream.messages == [{"role": "user", "content": "hi"}]
chunks: Final = stream.chunks
assert isinstance(chunks, list) and len(chunks) == 2

View file

@ -61,6 +61,7 @@ from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import prepare_response_for_header_attachment
def test_attach_guardrail_information_copies_recorded_entries_onto_model_response():
@ -7280,14 +7281,21 @@ class TestStreamingClientDisconnectBilling:
@pytest.mark.asyncio
async def test_disconnect_bills_partial_spend_for_anthropic_adapter_stream(self):
"""
/v1/messages wraps the chat stream in AnthropicStreamWrapper, which
hides the CustomStreamWrapper's collected chunks behind
.completion_stream; the partial-billing helper reads response.chunks,
so the wrapper must delegate inward or a disconnect bills nothing.
The proxy's cleanup gets the FallbackAwareAnthropicMessagesStream the
router returns for /v1/messages; its chunks/messages must delegate
through the translate_completion_output_params_streaming result to the
inner chat stream's collected chunks or a disconnect bills nothing.
"""
from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import (
AnthropicStreamWrapper,
AnthropicSSEStream,
)
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
AnthropicAdapter,
)
from litellm.router import FallbackAwareAnthropicMessagesStream
async def _sse_frames() -> AsyncGenerator[bytes, None]:
yield b"event: message_start\n\n"
recorder = _RecordingSuccessLogger()
original_callbacks = litellm.callbacks
@ -7295,14 +7303,20 @@ class TestStreamingClientDisconnectBilling:
try:
response = await self._start_partial_stream()
setattr(response.chunks[-1], "service_tier", "priority") # noqa: B010 # pydantic extra, not a declared field
wrapped: Final = AnthropicStreamWrapper(
completion_stream=response,
source_iterator: Final = AnthropicAdapter().translate_completion_output_params_streaming(
response,
model=response.model or "gpt-4o-mini",
is_async=True,
litellm_logging_obj=response.logging_obj,
)
assert isinstance(source_iterator, AnthropicSSEStream)
streamed: Final = prepare_response_for_header_attachment(
FallbackAwareAnthropicMessagesStream(_sse_frames(), source_iterator)
)
billed: Final = await _bill_partial_streamed_spend_on_disconnect(
{"litellm_logging_obj": response.logging_obj},
wrapped,
streamed,
)
for _ in range(50):