diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 5844c34eaa2..29053ae340b 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -902,7 +902,10 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): evt = self._events[self._idx] self._idx += 1 openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) in ( + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + ): self.completed_response = evt self._log_completed_response(is_async=True) return evt @@ -916,7 +919,10 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): evt = self._events[self._idx] self._idx += 1 openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) in ( + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + ): self.completed_response = evt self._log_completed_response(is_async=False) return evt @@ -969,7 +975,10 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): evt = self._events[self._idx] self._idx += 1 openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) in ( + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + ): self.completed_response = evt self._log_completed_response(is_async=True) return evt @@ -983,7 +992,10 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): evt = self._events[self._idx] self._idx += 1 openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) in ( + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + ): self.completed_response = evt self._log_completed_response(is_async=False) return evt diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_incomplete_status_transformation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_incomplete_status_transformation.py index 4ae7b8c75a5..f463864ed5d 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_incomplete_status_transformation.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_incomplete_status_transformation.py @@ -1,4 +1,6 @@ -from unittest.mock import AsyncMock +import asyncio + +from unittest.mock import AsyncMock, MagicMock import pytest @@ -8,7 +10,10 @@ from litellm.responses.litellm_completion_transformation.streaming_iterator impo from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, ) -from litellm.responses.streaming_iterator import _build_synthetic_response_events +from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + _build_synthetic_response_events, +) from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import Choices, Message, ModelResponse, Usage @@ -91,3 +96,18 @@ def test_replayed_stream_terminal_event_follows_status(finish_reason, event_type def test_omitted_temperature_defaults_to_zero(): assert _transform("stop", {}).temperature == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("finish_reason", ["length", "stop"]) +async def test_replayed_stream_logs_success_exactly_once(finish_reason): + logging_obj = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() + iterator = CachedResponsesAPIStreamingIterator( + response=_transform(finish_reason, {}), + logging_obj=logging_obj, + ) + async for _ in iterator: + pass + await asyncio.sleep(0) + assert logging_obj.dispatch_success_handlers.await_count == 1