fix(streaming): adapt provider errors to relay
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
@@ -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([
|
||||
|
||||
Reference in New Issue
Block a user