fix(router): unblock staging — mypy + coverage for aresponses streaming fallback (#28318)

Squash-merged by litellm-agent from cwang-otto's PR.
This commit is contained in:
cwang-otto 2026-05-20 12:04:48 +08:00 committed by Sameer Kankute
parent fa1214160a
commit ceeade246c
No known key found for this signature in database
2 changed files with 282 additions and 4 deletions

View file

@ -2296,8 +2296,13 @@ class Router:
from litellm.main import stream_chunk_builder
built = stream_chunk_builder(chunks=chunks)
if built is not None and built.usage is not None:
chat = built.usage
# stream_chunk_builder returns ModelResponse |
# TextCompletionResponse | None. ModelResponse sets .usage
# in __init__ rather than declaring it as a class field, so
# static narrowing doesn't expose it. Mirror the sync path
# (_completion_streaming_iterator) and pull via getattr.
chat = getattr(built, "usage", None) if built is not None else None
if chat is not None:
# getattr-with-default because the test path may
# substitute a SimpleNamespace lacking some fields;
# real Usage instances always have them.
@ -2388,7 +2393,12 @@ class Router:
and may regenerate same trade-off as the chat-completions path
for non-Anthropic fallbacks.
"""
base: List[Dict[str, Any]]
# base/continuation are List[Any] because ResponseInputParam items
# are a wide Union of TypedDicts (EasyInputMessageParam, Message,
# ResponseOutputMessageParam, ...) — annotating as List[Dict[str, Any]]
# rejects the list() spread of input_val. We cast the combined list to
# ResponseInputParam at the return.
base: List[Any]
if isinstance(input_val, str):
base = [
{
@ -2401,7 +2411,7 @@ class Router:
base = list(input_val)
else:
base = []
continuation: List[Dict[str, Any]] = [
continuation: List[Any] = [
{
"type": "message",
"role": "developer",

View file

@ -0,0 +1,268 @@
"""
Unit tests for the Responses-API streaming-fallback helpers added to Router
in PR #28215 (fix(router): wrap aresponses streaming iterator for mid-stream
fallbacks).
Targets the four helpers introduced on Router:
- _extract_partial_responses_usage
- _combine_responses_fallback_usage
- _build_responses_continuation_input
- _aresponses_streaming_iterator
"""
import os
import sys
from typing import Any, AsyncIterator, List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
from litellm import Router
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
def _make_router() -> Router:
return Router(
model_list=[
{
"model_name": "primary",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
},
},
{
"model_name": "fallback",
"litellm_params": {
"model": "openai/gpt-4o",
"api_key": "sk-test",
},
},
]
)
def _make_completed_event(
input_tokens: int, output_tokens: int, total_tokens: int
) -> ResponseCompletedEvent:
response = ResponsesAPIResponse.model_construct(
usage=ResponseAPIUsage(
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
)
)
return ResponseCompletedEvent.model_construct(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=response,
)
# -------- _extract_partial_responses_usage --------
def test_extract_partial_responses_usage_native_completed():
"""Native path: completed_response carries usage → returned as-is."""
completed = _make_completed_event(11, 7, 18)
source = MagicMock()
source.completed_response = completed
usage = Router._extract_partial_responses_usage(source)
assert usage is not None
assert usage.input_tokens == 11
assert usage.output_tokens == 7
assert usage.total_tokens == 18
def test_extract_partial_responses_usage_no_completed_response():
"""Native path: no completed_response → returns None."""
source = MagicMock()
source.completed_response = None
usage = Router._extract_partial_responses_usage(source)
assert usage is None
# -------- _combine_responses_fallback_usage --------
def test_combine_responses_fallback_usage_sums_completed_event():
"""Partial-stream usage is summed into the fallback event's usage."""
fallback_event = _make_completed_event(5, 3, 8)
partial = ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18)
Router._combine_responses_fallback_usage(fallback_event, partial)
combined = fallback_event.response.usage
assert combined is not None
assert combined.input_tokens == 16
assert combined.output_tokens == 10
assert combined.total_tokens == 26
def test_combine_responses_fallback_usage_passthrough_for_unknown_event():
"""Events that are not completed/failed/incomplete are not mutated."""
other = MagicMock() # not a ResponseCompletedEvent etc. → isinstance false
partial = ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2)
Router._combine_responses_fallback_usage(other, partial)
# No mutation expected on the unknown event — call is a no-op.
# -------- _build_responses_continuation_input --------
def test_build_responses_continuation_input_from_string():
out = Router._build_responses_continuation_input(
"Hello world", "partial assistant text"
)
assert len(out) == 3
assert out[0]["role"] == "user"
assert out[0]["content"][0]["text"] == "Hello world"
assert out[1]["role"] == "developer"
assert out[2]["role"] == "assistant"
assert out[2]["content"][0]["text"] == "partial assistant text"
def test_build_responses_continuation_input_from_list_preserves_items():
existing: List[Any] = [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "msg1"}],
}
]
out = Router._build_responses_continuation_input(existing, "partial")
assert len(out) == 3
assert out[0]["content"][0]["text"] == "msg1"
assert out[1]["role"] == "developer"
assert out[2]["role"] == "assistant"
def test_build_responses_continuation_input_from_none():
out = Router._build_responses_continuation_input(None, "partial")
assert len(out) == 2
assert out[0]["role"] == "developer"
assert out[1]["role"] == "assistant"
# -------- _aresponses_streaming_iterator (passthrough smoke test) --------
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_passthrough():
"""
Without MidStreamFallbackError, the wrapper yields source events
unchanged and returns a BaseResponsesAPIStreamingIterator subclass.
"""
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
events = [_make_completed_event(1, 1, 2)]
class _FakeSource:
"""Minimal source iterator. Provides every attribute the wrapper
constructor reads from source_iterator."""
def __init__(self) -> None:
self._i = 0
self.completed_response = None
self.response = MagicMock()
self.model = "openai/gpt-4o-mini"
self.logging_obj = MagicMock()
self.responses_api_provider_config = MagicMock()
self.start_time = 0.0
self.litellm_metadata = {}
self.custom_llm_provider = "openai"
self.request_data = {}
self.call_type = "aresponses"
self._hidden_params: dict = {}
def __aiter__(self) -> AsyncIterator[Any]:
return self
async def __anext__(self):
if self._i >= len(events):
raise StopAsyncIteration
ev = events[self._i]
self._i += 1
return ev
async def aclose(self):
return None
router = _make_router()
source = _FakeSource()
wrapper = await router._aresponses_streaming_iterator(
source, initial_kwargs={"model": "primary"}
)
assert isinstance(wrapper, BaseResponsesAPIStreamingIterator)
collected = [ev async for ev in wrapper]
assert len(collected) == 1
assert collected[0].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
# -------- _aresponses_with_streaming_fallbacks --------
@pytest.mark.asyncio
async def test_aresponses_with_streaming_fallbacks_non_streaming_passthrough():
"""Non-streaming response is returned unchanged, no wrap."""
router = _make_router()
plain_response = MagicMock()
async def fake_original(**_kwargs):
return plain_response
with patch.object(
router,
"_ageneric_api_call_with_fallbacks",
new=AsyncMock(return_value=plain_response),
):
out = await router._aresponses_with_streaming_fallbacks(
original_function=fake_original,
model="primary",
stream=False,
)
assert out is plain_response
@pytest.mark.asyncio
async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
"""Streaming response is wrapped via _aresponses_streaming_iterator."""
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
router = _make_router()
streaming_iter = MagicMock(spec=BaseResponsesAPIStreamingIterator)
wrapped = MagicMock(spec=BaseResponsesAPIStreamingIterator)
async def fake_original(**_kwargs):
return streaming_iter
with patch.object(
router,
"_ageneric_api_call_with_fallbacks",
new=AsyncMock(return_value=streaming_iter),
), patch.object(
router,
"_aresponses_streaming_iterator",
new=AsyncMock(return_value=wrapped),
) as mock_wrap:
out = await router._aresponses_with_streaming_fallbacks(
original_function=fake_original,
model="primary",
stream=True,
)
assert out is wrapped
mock_wrap.assert_awaited_once()