fix(api_server): emit final_response on chat-completions SSE when no deltas streamed (#31449)
Recovery paths (guardrail halt, partial_stream_recovery) can return a final_response without firing any content delta; /v1/chat/completions streaming then closed with an empty body. Mirror _ResponsesStream.collect_result. Co-authored-by: fmercurio <15571697+fmercurio@users.noreply.github.com>
This commit is contained in:
@@ -867,13 +867,14 @@ class OpenAICompatRoutesMixin:
|
||||
"""Stream ``chat.completion.chunk`` frames from the agent's delta queue. On client
|
||||
disconnect the agent is interrupted (stops LLM calls), then its task wrapper cancelled."""
|
||||
from gateway.platforms.api_server import (
|
||||
_abandon_agent_task, _chat_usage_payload, _sse_frame)
|
||||
_abandon_agent_task, _chat_usage_payload, _resolve_media_to_data_urls, _sse_frame)
|
||||
response = await self._prepare_sse_response(request, session_id, gateway_session_key)
|
||||
|
||||
def _chunk(delta: Dict[str, Any], finish_reason=None, **extra) -> Dict[str, Any]:
|
||||
return {"id": completion_id, "object": "chat.completion.chunk", "created": created,
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], **extra}
|
||||
content_sent = False
|
||||
try:
|
||||
await response.write(_sse_frame(_chunk({"role": "assistant"})))
|
||||
async for delta in _iter_stream_items(stream_q, agent_task, response):
|
||||
@@ -891,6 +892,8 @@ class OpenAICompatRoutesMixin:
|
||||
elif isinstance(delta, tuple) and len(delta) == 2 and delta[0] == "__approval__":
|
||||
await response.write(_sse_frame(delta[1], event="approval.request"))
|
||||
else:
|
||||
if delta:
|
||||
content_sent = True
|
||||
await response.write(_sse_frame(_chunk({"content": delta})))
|
||||
# The agent can fail after the queue drains (task raises / result flagged failed or
|
||||
# partial): surface a non-"stop" finish_reason like the non-streaming path.
|
||||
@@ -912,6 +915,14 @@ class OpenAICompatRoutesMixin:
|
||||
(isinstance(result, dict) and result.get("_notification_presentation_suppressed") is True)
|
||||
or getattr(agent_error, "_notification_presentation_suppressed", False) is True
|
||||
)
|
||||
# Recovery paths (guardrail halt, partial_stream_recovery, fallback prior-turn
|
||||
# content) can return a final_response without firing any content delta; emit it
|
||||
# once so the client does not see an empty stream (#31449). Mirrors
|
||||
# _ResponsesStream.collect_result for /v1/responses.
|
||||
if not content_sent and not presentation_muted and isinstance(result, dict):
|
||||
fallback_text = _resolve_media_to_data_urls(result.get("final_response") or "")
|
||||
if fallback_text:
|
||||
await response.write(_sse_frame(_chunk({"content": fallback_text})))
|
||||
if finish_reason != "stop":
|
||||
if err_msg and not presentation_muted:
|
||||
finish_chunk["error"] = {
|
||||
|
||||
51
tests/gateway/test_chat_completions_final_fallback.py
Normal file
51
tests/gateway/test_chat_completions_final_fallback.py
Normal file
@@ -0,0 +1,51 @@
|
||||
"""Streaming /v1/chat/completions must surface final_response when no content delta
|
||||
was streamed (guardrail halt / partial_stream_recovery paths, #31449)."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
def _run_stream(queued, final_response):
|
||||
from aiohttp import web
|
||||
from gateway.config import PlatformConfig
|
||||
from gateway.platforms.api_server import APIServerAdapter, ThreadSafeAsyncQueue
|
||||
|
||||
adapter = APIServerAdapter(PlatformConfig(enabled=True, token="test-key"))
|
||||
written = []
|
||||
|
||||
async def fake_agent():
|
||||
return {"final_response": final_response, "completed": True}, {
|
||||
"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}
|
||||
|
||||
async def run():
|
||||
stream_q = ThreadSafeAsyncQueue()
|
||||
for item in queued:
|
||||
stream_q.put_nowait(item)
|
||||
stream_q.put_nowait(None)
|
||||
agent_task = asyncio.ensure_future(fake_agent())
|
||||
resp = AsyncMock(spec=web.StreamResponse)
|
||||
resp.write = AsyncMock(side_effect=lambda data: written.append(data))
|
||||
resp.prepare = AsyncMock()
|
||||
req = MagicMock()
|
||||
req.headers = {}
|
||||
with patch("gateway.platforms.api_server.web.StreamResponse", return_value=resp):
|
||||
await adapter._write_sse_chat_completion(req, "cmpl-1", "m", 1, stream_q, agent_task)
|
||||
|
||||
asyncio.run(run())
|
||||
contents = []
|
||||
for frame in written:
|
||||
for line in frame.decode().splitlines():
|
||||
if line.startswith("data: {"):
|
||||
delta = json.loads(line[6:])["choices"][0]["delta"]
|
||||
if delta.get("content"):
|
||||
contents.append(delta["content"])
|
||||
return contents
|
||||
|
||||
|
||||
def test_final_response_emitted_when_no_deltas_streamed():
|
||||
assert _run_stream([], "recovered answer") == ["recovered answer"]
|
||||
|
||||
|
||||
def test_final_response_not_duplicated_after_streamed_deltas():
|
||||
assert _run_stream(["hello ", "world"], "hello world") == ["hello ", "world"]
|
||||
Reference in New Issue
Block a user