diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index 8dff2b596..0ebacd5dd 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -859,7 +859,7 @@ async def shield_violation_generator( api_params: ResponsesApiParams context: ResponsesContext Yields: - SSE-formatted strings for streaming events, ending with [DONE] + SSE-formatted strings for streaming events """ normalized_conv_id = normalize_conversation_id(api_params.conversation) available_quotas = get_available_quotas( @@ -935,8 +935,6 @@ async def shield_violation_generator( data_json = json.dumps(completed_event) yield f"event: response.completed\ndata: {data_json}\n\n" - yield "data: [DONE]\n\n" - def _sanitize_response_dict( response_dict: dict[str, Any], @@ -1249,8 +1247,6 @@ async def response_generator( latest_response_object.output, ) - yield "data: [DONE]\n\n" - async def generate_response( generator: AsyncIterator[str], @@ -1301,6 +1297,8 @@ async def generate_response( turn_summary.llm_response, ) _finalize_responses_root_span(root_span, turn_summary) + # Persist conversation state before clients can close the stream. + yield "data: [DONE]\n\n" finally: root_span.end() diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 38dc06d49..87aac2b2e 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -1228,7 +1228,7 @@ async def test_handle_streaming_blocked_returns_sse_consumes_shield_generator( assert "event: response.output_item.added" in body assert "event: response.output_item.done" in body assert "event: response.completed" in body - assert "[DONE]" in body + assert body.count("data: [DONE]") == 1 mock_client.responses.create.assert_not_called() @pytest.mark.asyncio