fix(streaming): adapt provider errors to relay

This commit is contained in:
beiyesi
2026-08-07 17:49:47 +08:00
committed by Teknium
parent 04bc5321c9
commit b55dd04711
2 changed files with 16 additions and 7 deletions
+11 -3
View File
@@ -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")
+5 -4
View File
@@ -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([