mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
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:
parent
fa1214160a
commit
ceeade246c
2 changed files with 282 additions and 4 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue