From b55dd047114ffd37624b04ee3dc8e148c5a4494e Mon Sep 17 00:00:00 2001 From: beiyesi Date: Fri, 7 Aug 2026 17:49:47 +0800 Subject: [PATCH] fix(streaming): adapt provider errors to relay --- agent/chat_completion_helpers.py | 14 +++++++++++--- tests/run_agent/test_run_agent.py | 9 +++++---- 2 files changed, 16 insertions(+), 7 deletions(-) diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index ea583c483d..3c5a90370a 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -299,14 +299,17 @@ def _provider_stream_error_from_json_decode_error( ) -def _iter_provider_stream_chunks(stream): +def _iter_provider_stream_chunks(stream, *, response: Any = None): """Yield SDK chunks while translating SDK-level SSE decode failures.""" try: yield from stream except json.JSONDecodeError as error: + stream_response = response() if callable(response) else response + if stream_response is None: + stream_response = getattr(stream, "response", None) raise _provider_stream_error_from_json_decode_error( error, - response=getattr(stream, "response", None), + response=stream_response, ) from error @@ -3717,6 +3720,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= request_client_holder["diag"] = _diag _writer_token = {"value": None} attempt_request_client = {"value": None} + attempt_stream_response = {"value": None} def _open_stream(next_api_kwargs: dict[str, Any]): stream_kwargs = { @@ -3745,6 +3749,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= def _stream_created(raw_stream: Any) -> None: response = getattr(raw_stream, "response", None) + attempt_stream_response["value"] = response agent._capture_rate_limits(response) agent._capture_credits(response) agent._stream_diag_capture_response(_diag, response) @@ -3852,7 +3857,10 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= except Exception: pass - for chunk in _iter_provider_stream_chunks(stream): + for chunk in _iter_provider_stream_chunks( + stream, + response=lambda: attempt_stream_response["value"], + ): last_chunk_time["t"] = time.time() agent._touch_activity("receiving stream response") diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index c6370ba6fd..7ad5c391b0 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -5796,10 +5796,11 @@ class TestStreamingApiCall: resp = agent._interruptible_streaming_api_call({"messages": []}) assert resp.choices[0].message.content == error_text - assert resp.choices[0].finish_reason == "stop" - assert [ - call.args[0] for call in agent.stream_delta_callback.call_args_list - ] == [error_text] + # Current main treats every text-only stream without a terminal finish + # signal as a partial response. The SSE-shaped text remains literal, + # but is withheld from the callback so the retry path can own delivery. + assert resp.choices[0].finish_reason == "length" + agent.stream_delta_callback.assert_not_called() def test_run_conversation_retries_stream_error_finish_rate_limit(self, agent): first_attempt = iter([