diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index 58fef45a49a..8b2dd47bc6b 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -2366,6 +2366,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= pass def _bedrock_call(): + stream = None try: from agent import relay_llm from agent.bedrock_adapter import ( @@ -2476,6 +2477,9 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= result["response"] = stream.final_response or streamed_response except Exception as e: result["error"] = e + finally: + if stream is not None: + stream.close() t = threading.Thread( target=_context_thread_target(_bedrock_call), daemon=True @@ -3029,6 +3033,8 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= if hasattr(chunk, "usage") and chunk.usage: usage_obj = chunk.usage + stream.close() + if _stream_attempt_was_cancelled(stream_attempt_id): raise _httpx.RemoteProtocolError( f"stream attempt {stream_attempt_id} was superseded" @@ -3351,9 +3357,12 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= ) from None raise finally: - manager = _stream_context["manager"] - if manager is not None: - manager.__exit__(None, None, None) + try: + stream.close() + finally: + manager = _stream_context["manager"] + if manager is not None: + manager.__exit__(None, None, None) if agent._interrupt_requested: return None diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 51338bdf398..8b799b6e775 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -304,6 +304,7 @@ class ManagedLlmStream(Iterator[Any]): self.final_response: Any = None self._loop: asyncio.AbstractEventLoop | None = None self._stream: Any = None + self._raw_stream_resource: Any = None self._closed = False self._callback_error: BaseException | None = None self._logical: tuple[relay_runtime.RelayTurnContext, Any, str] | None = None @@ -329,6 +330,7 @@ class ManagedLlmStream(Iterator[Any]): self.final_response = raw_stream self._stream = iter(()) else: + self._raw_stream_resource = raw_stream if on_stream_created is not None: on_stream_created(raw_stream) self._stream = iter(raw_stream) @@ -497,9 +499,23 @@ class ManagedLlmStream(Iterator[Any]): loop = self._loop self._loop = None if loop is None: - close = getattr(self._stream, "close", None) - if callable(close): - close() + resources = (self._stream, self._raw_stream_resource) + self._stream = None + self._raw_stream_resource = None + closed_ids: set[int] = set() + for resource in resources: + if resource is None or id(resource) in closed_ids: + continue + closed_ids.add(id(resource)) + close = getattr(resource, "close", None) + if callable(close): + try: + close() + except Exception: + logger.debug( + "Provider stream cleanup failed", + exc_info=True, + ) if not self._defer_logical_completion: _complete_logical(self._logical, outcome=logical_outcome) self._logical = None diff --git a/tests/agent/test_bedrock_interrupt_post_worker.py b/tests/agent/test_bedrock_interrupt_post_worker.py index 0ee5e4fce33..52c6877d1f4 100644 --- a/tests/agent/test_bedrock_interrupt_post_worker.py +++ b/tests/agent/test_bedrock_interrupt_post_worker.py @@ -86,7 +86,21 @@ def test_bedrock_stream_returns_normally_when_not_interrupted(): agent._interrupt_requested = False resp = SimpleNamespace(choices=[], usage=None, stop_reason="end_turn") - fake_client = SimpleNamespace(converse_stream=lambda **kw: {"stream": []}) + + class ProviderStream: + def __init__(self): + self.closed = False + + def __iter__(self): + return iter(()) + + def close(self): + self.closed = True + + provider_stream = ProviderStream() + fake_client = SimpleNamespace( + converse_stream=lambda **kw: {"stream": provider_stream} + ) with patch("agent.bedrock_adapter._get_bedrock_runtime_client", return_value=fake_client), \ patch("agent.bedrock_adapter.stream_converse_with_callbacks", return_value=resp), \ @@ -97,3 +111,4 @@ def test_bedrock_stream_returns_normally_when_not_interrupted(): api_kwargs = {"__bedrock_region__": "us-east-1", "__bedrock_converse__": True} out = cch.interruptible_streaming_api_call(agent, api_kwargs) assert out is resp + assert provider_stream.closed is True diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index 274fc351647..75367fb93a1 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -214,6 +214,39 @@ def test_non_deferred_partial_stream_close_cancels_logical_call( assert {"outcome": "cancelled"} in terminal_outputs +def test_direct_stream_close_reaches_original_provider_resource(monkeypatch): + class ProviderStream: + def __init__(self): + self.closed = False + + def __iter__(self): + return iter([{"delta": "partial"}, {"delta": "unused"}]) + + def close(self): + self.closed = True + + provider_stream = ProviderStream() + monkeypatch.setattr( + relay_runtime, + "resolve_execution_context", + lambda _session_id: (None, None, None), + ) + + stream = relay_llm.stream( + {"model": "test-model", "messages": []}, + lambda _request: provider_stream, + session_id="session-1", + name="test-provider", + model_name="test-model", + finalizer=dict, + ) + + assert next(stream) == {"delta": "partial"} + stream.close() + + assert provider_stream.closed is True + + def test_non_stream_preserves_raw_provider_response_identity(relay_turn): _relay, _turn = relay_turn raw_response = SimpleNamespace(model="test-model", content="raw") diff --git a/tests/run_agent/test_streaming.py b/tests/run_agent/test_streaming.py index 55a0b9c008e..ceebff03525 100644 --- a/tests/run_agent/test_streaming.py +++ b/tests/run_agent/test_streaming.py @@ -95,6 +95,51 @@ class TestStreamingAccumulator: assert response.usage is not None assert response.usage.completion_tokens == 3 + @patch("run_agent.AIAgent._create_request_openai_client") + @patch("run_agent.AIAgent._close_request_openai_client") + def test_chat_stream_closes_original_provider_resource( + self, + mock_close, + mock_create, + ): + from run_agent import AIAgent + + class ProviderStream: + def __init__(self): + self.closed = False + + def __iter__(self): + return iter([ + _make_stream_chunk( + content="Hello", + finish_reason="stop", + model="test-model", + ) + ]) + + def close(self): + self.closed = True + + provider_stream = ProviderStream() + mock_client = MagicMock() + mock_client.chat.completions.create.return_value = provider_stream + mock_create.return_value = mock_client + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1", + model="test/model", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + agent.api_mode = "chat_completions" + agent._interrupt_requested = False + + response = agent._interruptible_streaming_api_call({}) + + assert response.choices[0].message.content == "Hello" + assert provider_stream.closed is True + @patch("run_agent.AIAgent._create_request_openai_client") @patch("run_agent.AIAgent._close_request_openai_client") def test_native_gemini_endpoint_omits_stream_options(self, mock_close, mock_create): @@ -1203,6 +1248,7 @@ class TestAnthropicStreamCallbacks: agent._interruptible_streaming_api_call({}) assert touch_calls.count("receiving stream response") == len(events) + mock_stream.close.assert_called_once() @patch("run_agent.AIAgent._rebuild_anthropic_client") @patch("run_agent.AIAgent._replace_primary_openai_client")