diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index f270442d..bc7244e1 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1426,6 +1426,14 @@ class BaseUpstreamProvider: if done_seen: yield b"data: [DONE]\n\n" + except httpx.RemoteProtocolError as stream_error: + logger.warning( + "Upstream stream ended before the response was complete", + extra={ + "error": str(stream_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) except Exception as stream_error: logger.warning( "Streaming interrupted; finalizing before closing upstream", @@ -1869,6 +1877,14 @@ class BaseUpstreamProvider: if done_seen: yield b"data: [DONE]\n\n" + except httpx.RemoteProtocolError as stream_error: + logger.warning( + "Upstream Responses API stream ended before the response was complete", + extra={ + "error": str(stream_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) except Exception as stream_error: logger.warning( "Responses API streaming interrupted; finalizing before closing upstream", diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index e70804c1..d60b598a 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -404,11 +404,8 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( client=client, ) emitted = bytearray() - with pytest.raises(httpx.RemoteProtocolError): - async for chunk in response.body_iterator: - emitted.extend( - chunk.encode() if isinstance(chunk, str) else bytes(chunk) - ) + async for chunk in response.body_iterator: + emitted.extend(chunk.encode() if isinstance(chunk, str) else bytes(chunk)) adjust.assert_awaited_once() if finalization_fails: @@ -423,7 +420,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( @pytest.mark.asyncio @pytest.mark.parametrize("api", ["chat", "responses"]) -async def test_partial_stream_preserves_transport_error_when_billing_db_is_down( +async def test_partial_stream_closes_when_billing_db_is_down( api: str, ) -> None: provider = BaseUpstreamProvider( @@ -476,9 +473,8 @@ async def test_partial_stream_preserves_transport_error_when_billing_db_is_down( reservation_snapshot=snapshot, client=client, ) - with pytest.raises(httpx.RemoteProtocolError, match="incomplete chunked read"): - async for _ in response.body_iterator: - pass + async for _ in response.body_iterator: + pass upstream_response.aclose.assert_awaited_once() client.aclose.assert_awaited_once()