mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-31 19:16:29 +00:00
fix(relay): close managed provider streams
Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
parent
f370af6816
commit
c789bdd38a
5 changed files with 126 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue