diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index 505b5b09433..7f6b7efb955 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -13,7 +13,11 @@ from litellm.responses.litellm_completion_transformation.streaming_iterator impo from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, ) -from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator +from litellm.responses.streaming_iterator import ( + BaseResponsesAPIStreamingIterator, + MockResponsesAPIStreamingIterator, +) +from litellm.types.integrations.custom_logger import converted_stream_requested from litellm.types.llms.openai import ( ResponseInputParam, ResponsesAPIOptionalRequestParams, @@ -81,6 +85,18 @@ class LiteLLMCompletionTransformationHandler: ) ) + converted_stream: Final = converted_stream_requested(kwargs) or converted_stream_requested( + litellm_completion_request + ) + if converted_stream and not kwargs.get("_agentic_loop_depth"): + return MockResponsesAPIStreamingIterator( + model=model, + transformed_response=responses_api_response, + logging_obj=kwargs.get("logging_obj"), + custom_llm_provider=custom_llm_provider, + litellm_metadata=kwargs.get("litellm_metadata"), + ) + return responses_api_response elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper): @@ -126,6 +142,18 @@ class LiteLLMCompletionTransformationHandler: ) ) + converted_stream: Final = converted_stream_requested(kwargs) or converted_stream_requested( + litellm_completion_request + ) + if converted_stream and not kwargs.get("_agentic_loop_depth"): + return MockResponsesAPIStreamingIterator( + model=litellm_completion_request.get("model") or "", + transformed_response=responses_api_response, + logging_obj=kwargs.get("logging_obj"), + custom_llm_provider=litellm_completion_request.get("custom_llm_provider"), + litellm_metadata=kwargs.get("litellm_metadata"), + ) + return responses_api_response elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper): diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 5e045c3e84f..d0d8c203d9d 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -334,9 +334,15 @@ class BaseResponsesAPIStreamingIterator: # set hidden params for response headers (e.g., x-litellm-model-id) # This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py + _model_call_details: Final = getattr(self.logging_obj, "model_call_details", None) + _optional_params: Final = ( + _typed_gets_litellm_params(_model_call_details.get)("litellm_params", {}) # mutable-ok: fallback dict + if isinstance(_model_call_details, dict) + else {} # mutable-ok: fallback empty mapping + ) _api_base: Final = get_api_base( model=model or "", - optional_params=_typed_gets_litellm_params(self.logging_obj.model_call_details.get)("litellm_params", {}), + optional_params=_optional_params, ) self._hidden_params: dict[str, object] = { "model_id": _model_id_from_metadata(litellm_metadata), @@ -532,7 +538,7 @@ class BaseResponsesAPIStreamingIterator: raise def _log_completed_response(self, *, is_async: bool) -> None: - if self._completed_response_logged: + if self._completed_response_logged or self.logging_obj is None: return self._completed_response_logged = True @@ -1191,22 +1197,35 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __init__( self, - response: httpx.Response, - model: str, - responses_api_provider_config: BaseResponsesAPIConfig, - logging_obj: LiteLLMLoggingObj, + response: httpx.Response | None = None, + model: str = "", + responses_api_provider_config: BaseResponsesAPIConfig | None = None, + logging_obj: LiteLLMLoggingObj | None = None, litellm_metadata: dict[str, object] | None = None, custom_llm_provider: str | None = None, request_data: dict[str, object] | None = None, call_type: str | None = None, + transformed_response: ResponsesAPIResponse | None = None, ): - transformed: Final = responses_api_provider_config.transform_response_api_response( - model=model, - raw_response=response, - logging_obj=logging_obj, + transformed: Final[ResponsesAPIResponse | None] = ( + transformed_response + if transformed_response is not None + else ( + responses_api_provider_config.transform_response_api_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + if responses_api_provider_config is not None and response is not None + else None + ) ) + if transformed is None: + raise ValueError( + "Either transformed_response or both responses_api_provider_config and response must be provided to MockResponsesAPIStreamingIterator" + ) super().__init__( - response=httpx.Response(200), + response=response or httpx.Response(200), model=model, responses_api_provider_config=None, logging_obj=logging_obj, diff --git a/tests/unit/responses/litellm_completion_transformation/test_handler.py b/tests/unit/responses/litellm_completion_transformation/test_handler.py index bb374b90f4e..d8633f59ac0 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_handler.py +++ b/tests/unit/responses/litellm_completion_transformation/test_handler.py @@ -223,3 +223,136 @@ async def test_bridged_follow_up_turn_keeps_the_addressed_response_id_off_the_pr ) assert isinstance(response, ResponsesAPIResponse) assert [item.type for item in response.output] == ["message"] + + +@pytest.mark.parametrize( + "converted_stream_flag", + [ + "_websearch_interception_converted_stream", + "_code_interpreter_interception_converted_stream", + "_headroom_interception_converted_stream", + ], +) +@pytest.mark.asyncio +async def test_async_fallback_wraps_converted_stream_as_synthetic_stream(converted_stream_flag): + from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator + from litellm.types.utils import Choices, Message, ModelResponse + + handler = LiteLLMCompletionTransformationHandler() + + async def fake_acompletion(**kwargs): + return ModelResponse( + id="chatcmpl-test", + created=1, + model="gpt-4o", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Search summary", role="assistant"), + ) + ], + ) + + with patch("litellm.acompletion", fake_acompletion): + response = await handler.response_api_handler( + model="gpt-4o", + input="Search for the weather", + responses_api_request={}, + custom_llm_provider="hosted_vllm", + _is_async=True, + **{converted_stream_flag: True}, + ) + + assert isinstance(response, BaseResponsesAPIStreamingIterator) + events = [event async for event in response] + assert len(events) > 0 + assert getattr(events[-1], "type", None) == "response.completed" + + +@pytest.mark.parametrize( + "converted_stream_flag", + [ + "_websearch_interception_converted_stream", + "_code_interpreter_interception_converted_stream", + "_headroom_interception_converted_stream", + ], +) +def test_sync_fallback_wraps_converted_stream_as_synthetic_stream(converted_stream_flag): + from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator + from litellm.types.utils import Choices, Message, ModelResponse + + handler = LiteLLMCompletionTransformationHandler() + + def fake_completion(**kwargs): + return ModelResponse( + id="chatcmpl-test-sync", + created=1, + model="gpt-4o", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Sync search summary", role="assistant"), + ) + ], + ) + + with patch("litellm.completion", fake_completion): + response = handler.response_api_handler( + model="gpt-4o", + input="Search for the weather", + responses_api_request={}, + custom_llm_provider="hosted_vllm", + _is_async=False, + **{converted_stream_flag: True}, + ) + + assert isinstance(response, BaseResponsesAPIStreamingIterator) + events = list(response) + assert len(events) > 0 + assert getattr(events[-1], "type", None) == "response.completed" + + +def test_mock_responses_streaming_iterator_validation_and_config_branches(): + import httpx + + from litellm.responses.streaming_iterator import MockResponsesAPIStreamingIterator + from litellm.types.llms.openai import ResponsesAPIResponse + + with pytest.raises(ValueError, match="Either transformed_response or both"): + MockResponsesAPIStreamingIterator() + + class _MockConfig: + def transform_response_api_response(self, **kwargs): + return ResponsesAPIResponse( + id="resp_cfg_test", + created_at=1, + status="completed", + model="test-model", + object="response", + output=[], + ) + + class _MockLoggingObj: + def __init__(self): + self.model_call_details = {"litellm_params": {"api_key": "fake"}} + + async def async_success_handler(self, *args, **kwargs): + pass + + def success_handler(self, *args, **kwargs): + pass + + logging_obj = _MockLoggingObj() + iterator = MockResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="gpt-4o", + responses_api_provider_config=_MockConfig(), + logging_obj=logging_obj, + ) + events = list(iterator) + assert len(events) > 0 + assert getattr(events[-1], "type", None) == "response.completed"