diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index 3c89ada157..796cc53a4f 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -3530,6 +3530,12 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): headers = { "Content-Type": "text/event-stream", "Cache-Control": "no-cache", "X-Accel-Buffering": "no", **self._session_headers(session_id, gateway_session_key)} + # CORS middleware can't inject headers into StreamResponse after + # prepare() flushes them, so resolve CORS headers up front (#72892). + origin = request.headers.get("Origin", "") + cors = self._cors_headers_for_origin(origin) if origin else None + if cors: + headers.update(cors) response = web.StreamResponse(status=200, headers=headers) await response.prepare(request) try: diff --git a/tests/gateway/test_api_server.py b/tests/gateway/test_api_server.py index 9b190e89b5..ec007ce696 100644 --- a/tests/gateway/test_api_server.py +++ b/tests/gateway/test_api_server.py @@ -2263,6 +2263,61 @@ class TestCORS: assert "Authorization" in resp.headers.get("Access-Control-Allow-Headers", "") + @pytest.mark.asyncio + async def test_cors_headers_present_on_session_chat_stream(self): + """SSE streaming endpoint must include CORS headers on the response. + + The CORS middleware cannot inject headers into a ``StreamResponse`` + after ``prepare()`` flushes them, so the handler must resolve CORS + headers up front — same pattern as ``/v1/chat/completions`` and + ``/v1/responses`` streaming paths (#72892). + """ + adapter = _make_adapter(cors_origins=["http://localhost:3000"]) + app = _create_app(adapter) + async with TestClient(TestServer(app)) as cli: + with ( + patch.object(adapter, "_get_existing_session_or_404", return_value=({"id": "s1"}, None)), + patch.object(adapter, "_conversation_history_for_session", return_value=[]), + patch.object(adapter, "_run_agent", new_callable=AsyncMock) as mock_run, + ): + mock_run.return_value = ( + {"final_response": "ok", "messages": [], "api_calls": 1}, + {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + ) + resp = await cli.post( + "/api/sessions/s1/chat/stream", + json={"message": "hi"}, + headers={"Origin": "http://localhost:3000"}, + ) + assert resp.status == 200 + await resp.text() # consume SSE stream fully + assert resp.headers.get("Access-Control-Allow-Origin") == "http://localhost:3000" + assert "POST" in resp.headers.get("Access-Control-Allow-Methods", "") + + @pytest.mark.asyncio + async def test_cors_headers_absent_on_session_chat_stream_no_origin(self): + """No CORS headers on the streaming endpoint when no Origin is sent.""" + adapter = _make_adapter() + app = _create_app(adapter) + async with TestClient(TestServer(app)) as cli: + with ( + patch.object(adapter, "_get_existing_session_or_404", return_value=({"id": "s1"}, None)), + patch.object(adapter, "_conversation_history_for_session", return_value=[]), + patch.object(adapter, "_run_agent", new_callable=AsyncMock) as mock_run, + ): + mock_run.return_value = ( + {"final_response": "ok", "messages": [], "api_calls": 1}, + {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + ) + resp = await cli.post( + "/api/sessions/s1/chat/stream", + json={"message": "hi"}, + ) + assert resp.status == 200 + await resp.text() + assert resp.headers.get("Access-Control-Allow-Origin") is None + + # --------------------------------------------------------------------------- # Conversation parameter # ---------------------------------------------------------------------------