fix(relay): close managed provider streams

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-23 13:32:03 -07:00
parent f370af6816
commit c789bdd38a
5 changed files with 126 additions and 7 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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")