mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge 658b778109 into af6dc1db08
This commit is contained in:
commit
169bca6d1d
2 changed files with 72 additions and 21 deletions
|
|
@ -93,6 +93,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self.responses_api_request: ResponsesAPIOptionalRequestParams = responses_api_request
|
||||
self.custom_llm_provider: str | None = custom_llm_provider
|
||||
self.litellm_metadata: dict | None = litellm_metadata or {}
|
||||
_hp: Final = getattr(litellm_custom_stream_wrapper, "_hidden_params", None)
|
||||
self._hidden_params: dict[str, object] = dict(_hp) if isinstance(_hp, dict) else {} # mutable-ok: router header
|
||||
# Store lightweight dict snapshots for stream_chunk_builder to reduce
|
||||
# repeated Pydantic attribute access in end-of-stream assembly.
|
||||
self.collected_chat_completion_chunks: list[dict[str, object]] = []
|
||||
|
|
|
|||
|
|
@ -10,12 +10,12 @@ Targets the four helpers introduced on Router:
|
|||
- _aresponses_streaming_iterator
|
||||
"""
|
||||
|
||||
from typing import Any, AsyncIterator, List
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm import Router
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseAPIUsage,
|
||||
|
|
@ -46,9 +46,7 @@ def _make_router() -> Router:
|
|||
)
|
||||
|
||||
|
||||
def _make_completed_event(
|
||||
input_tokens: int, output_tokens: int, total_tokens: int
|
||||
) -> ResponseCompletedEvent:
|
||||
def _make_completed_event(input_tokens: int, output_tokens: int, total_tokens: int) -> ResponseCompletedEvent:
|
||||
response = ResponsesAPIResponse.model_construct(
|
||||
usage=ResponseAPIUsage(
|
||||
input_tokens=input_tokens,
|
||||
|
|
@ -145,9 +143,7 @@ def test_combine_responses_fallback_usage_passthrough_for_unknown_event():
|
|||
|
||||
|
||||
def test_build_responses_continuation_input_from_string():
|
||||
out = Router._build_responses_continuation_input(
|
||||
"Hello world", "partial assistant text"
|
||||
)
|
||||
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"
|
||||
|
|
@ -157,7 +153,7 @@ def test_build_responses_continuation_input_from_string():
|
|||
|
||||
|
||||
def test_build_responses_continuation_input_from_list_preserves_items():
|
||||
existing: List[Any] = [
|
||||
existing: list[Any] = [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
|
|
@ -227,9 +223,7 @@ async def test_aresponses_streaming_iterator_passthrough():
|
|||
router = _make_router()
|
||||
source = _FakeSource()
|
||||
|
||||
wrapper = await router._aresponses_streaming_iterator(
|
||||
source, initial_kwargs={"model": "primary"}
|
||||
)
|
||||
wrapper = await router._aresponses_streaming_iterator(source, initial_kwargs={"model": "primary"})
|
||||
assert isinstance(wrapper, BaseResponsesAPIStreamingIterator)
|
||||
|
||||
collected = [ev async for ev in wrapper]
|
||||
|
|
@ -276,15 +270,18 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
|
|||
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:
|
||||
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",
|
||||
|
|
@ -509,3 +506,55 @@ async def test_aresponses_client_error_event_skips_fallback():
|
|||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_fallback.assert_not_awaited()
|
||||
|
||||
|
||||
# -------- regression: real bridge iterator honors the base contract (LIT-4912) --------
|
||||
|
||||
|
||||
def test_extract_partial_responses_usage_real_bridge_iterator_pre_first_chunk():
|
||||
"""
|
||||
Regression for LIT-4912.
|
||||
|
||||
When a provider errors before the first content chunk, the completion-bridge
|
||||
iterator reaches _extract_partial_responses_usage with no collected chunks.
|
||||
A real LiteLLMCompletionStreamingIterator (not a MagicMock standing in for
|
||||
it) must expose the base-class contract so usage extraction returns None
|
||||
instead of raising AttributeError on completed_response, and must carry
|
||||
_hidden_params so the router does not pre-wrap it for header attachment and
|
||||
thereby skip the mid-stream fallback path entirely.
|
||||
"""
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
class _FakeStreamWrapper:
|
||||
def __init__(self) -> None:
|
||||
self.logging_obj = MagicMock()
|
||||
self._hidden_params = {"model_id": "deployment-123"}
|
||||
|
||||
class _FakeStreamWrapperWithoutHiddenParams:
|
||||
def __init__(self) -> None:
|
||||
self.logging_obj = MagicMock()
|
||||
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
litellm_custom_stream_wrapper=_FakeStreamWrapper(),
|
||||
request_input="hi",
|
||||
responses_api_request={},
|
||||
)
|
||||
|
||||
assert isinstance(iterator, BaseResponsesAPIStreamingIterator)
|
||||
assert iterator.completed_response is None
|
||||
assert iterator._hidden_params == {"model_id": "deployment-123"}
|
||||
assert not iterator.collected_chat_completion_chunks
|
||||
|
||||
iterator_without_hidden_params = LiteLLMCompletionStreamingIterator(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
litellm_custom_stream_wrapper=_FakeStreamWrapperWithoutHiddenParams(),
|
||||
request_input="hi",
|
||||
responses_api_request={},
|
||||
)
|
||||
|
||||
assert iterator_without_hidden_params._hidden_params == {}
|
||||
assert Router._extract_partial_responses_usage(iterator) is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue