fix(agent): bound sequential tool calls
This commit is contained in:
@@ -387,6 +387,10 @@ class _ManagedToolResult:
|
||||
dispatched: bool
|
||||
|
||||
|
||||
class _ToolTimeoutResult(str):
|
||||
"""Marker for a synthesized sequential-tool timeout result."""
|
||||
|
||||
|
||||
class _ConcurrentToolAuthorizationGate:
|
||||
"""Serialize policy prompts and exclude human approval waits from batch deadlines.
|
||||
|
||||
@@ -661,6 +665,117 @@ def _run_agent_tool_execution_middleware(
|
||||
)
|
||||
|
||||
|
||||
def _run_sequential_tool_execution_middleware(
|
||||
agent,
|
||||
*,
|
||||
function_name: str,
|
||||
function_args: dict,
|
||||
effective_task_id: str,
|
||||
tool_call_id: str,
|
||||
execute,
|
||||
scope_block: str | None = None,
|
||||
display_index: int | None = None,
|
||||
middleware_trace: list[dict[str, Any]] | None = None,
|
||||
) -> _ManagedToolResult:
|
||||
"""Run one sequential call with the concurrent executor's deadline."""
|
||||
timeout_s = _resolve_concurrent_tool_timeout()
|
||||
kwargs = {
|
||||
"function_name": function_name,
|
||||
"function_args": function_args,
|
||||
"effective_task_id": effective_task_id,
|
||||
"tool_call_id": tool_call_id,
|
||||
"execute": execute,
|
||||
"scope_block": scope_block,
|
||||
"display_index": display_index,
|
||||
"middleware_trace": middleware_trace,
|
||||
}
|
||||
if timeout_s is None:
|
||||
return _run_agent_tool_execution_middleware(agent, **kwargs)
|
||||
|
||||
from tools.daemon_pool import DaemonThreadPoolExecutor
|
||||
|
||||
authorization_gate = _ConcurrentToolAuthorizationGate()
|
||||
worker_tid: list[int] = []
|
||||
|
||||
def _run() -> _ManagedToolResult:
|
||||
tid = threading.current_thread().ident
|
||||
worker_tid.append(tid)
|
||||
with agent._tool_worker_threads_lock:
|
||||
agent._tool_worker_threads.add(tid)
|
||||
try:
|
||||
return _run_agent_tool_execution_middleware(
|
||||
agent, authorization_gate=authorization_gate, **kwargs
|
||||
)
|
||||
finally:
|
||||
with agent._tool_worker_threads_lock:
|
||||
agent._tool_worker_threads.discard(tid)
|
||||
try:
|
||||
_ra()._set_interrupt(False, tid)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
executor = DaemonThreadPoolExecutor(max_workers=1)
|
||||
future = executor.submit(propagate_context_to_thread(_run))
|
||||
deadline = time.monotonic() + timeout_s
|
||||
started = time.monotonic()
|
||||
timed_out = False
|
||||
try:
|
||||
while True:
|
||||
remaining = (
|
||||
deadline + authorization_gate.excluded_seconds() - time.monotonic()
|
||||
)
|
||||
if remaining <= 0:
|
||||
timed_out = True
|
||||
break
|
||||
try:
|
||||
return future.result(timeout=min(5.0, remaining))
|
||||
except concurrent.futures.TimeoutError:
|
||||
elapsed = int(time.monotonic() - started)
|
||||
if elapsed > 0 and elapsed % 30 < 5:
|
||||
agent._touch_activity(
|
||||
f"sequential tool running ({elapsed}s): {function_name}"
|
||||
)
|
||||
|
||||
message = (
|
||||
f"Error executing tool '{function_name}': "
|
||||
f"timed out after {timeout_s:.1f}s"
|
||||
)
|
||||
logger.warning(
|
||||
"sequential tool %s timed out after %.1fs", function_name, timeout_s
|
||||
)
|
||||
future.cancel()
|
||||
for tid in worker_tid:
|
||||
try:
|
||||
_ra()._set_interrupt(True, tid)
|
||||
except Exception:
|
||||
pass
|
||||
trace = middleware_trace if middleware_trace is not None else []
|
||||
_emit_terminal_post_tool_call(
|
||||
agent,
|
||||
function_name=function_name,
|
||||
function_args=function_args,
|
||||
result=message,
|
||||
effective_task_id=effective_task_id,
|
||||
tool_call_id=tool_call_id,
|
||||
duration_ms=int(timeout_s * 1000),
|
||||
status="timeout",
|
||||
error_type="tool_timeout",
|
||||
error_message=message,
|
||||
middleware_trace=list(trace),
|
||||
)
|
||||
return _ManagedToolResult(
|
||||
result=_ToolTimeoutResult(message),
|
||||
args=function_args,
|
||||
middleware_trace=trace,
|
||||
blocked=False,
|
||||
dispatched=True,
|
||||
)
|
||||
finally:
|
||||
# Never join a wedged worker. DaemonThreadPoolExecutor also keeps it out
|
||||
# of the stdlib atexit join, matching the concurrent timeout path.
|
||||
executor.shutdown(wait=not timed_out, cancel_futures=timed_out)
|
||||
|
||||
|
||||
def _begin_tool_execution(
|
||||
agent,
|
||||
*,
|
||||
@@ -1608,6 +1723,12 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
"""
|
||||
# Resolve the context-scaled tool-output budget once per turn.
|
||||
_tool_budget = _budget_for_agent(agent)
|
||||
|
||||
# Keep every runtime-tool branch on one bounded execution funnel without
|
||||
# duplicating timeout policy across the branch-specific callbacks below.
|
||||
def _run_agent_tool_execution_middleware(agent, **kwargs):
|
||||
return _run_sequential_tool_execution_middleware(agent, **kwargs)
|
||||
|
||||
for i, tool_call in enumerate(assistant_message.tool_calls, 1):
|
||||
if getattr(agent, "_incremental_persistence_failed", False):
|
||||
return
|
||||
@@ -2193,6 +2314,7 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
logger.error("handle_function_call raised for %s: %s", function_name, tool_error, exc_info=True)
|
||||
tool_duration = time.time() - tool_start_time
|
||||
|
||||
_execution_timed_out = isinstance(function_result, _ToolTimeoutResult)
|
||||
if isinstance(function_result, str):
|
||||
result_preview = function_result if agent.verbose_logging else (
|
||||
function_result[:200] if len(function_result) > 200 else function_result
|
||||
@@ -2215,6 +2337,7 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
from agent.agent_runtime_helpers import agent_runtime_owns_post_tool_hook
|
||||
_executor_must_emit_post_hook = (
|
||||
not _execution_blocked
|
||||
and not _execution_timed_out
|
||||
and (
|
||||
not _execution_dispatched
|
||||
or agent_runtime_owns_post_tool_hook(agent, function_name)
|
||||
@@ -2287,7 +2410,12 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
# Unwrap _multimodal dicts to an OpenAI-style content list
|
||||
# (see parallel path for rationale). String results pass through.
|
||||
_tool_content = agent._tool_result_content_for_active_model(function_name, function_result)
|
||||
tool_message = make_tool_result_message(function_name, _tool_content, tool_call.id)
|
||||
tool_message = make_tool_result_message(
|
||||
function_name,
|
||||
_tool_content,
|
||||
tool_call.id,
|
||||
effect_disposition="unknown" if _execution_timed_out else None,
|
||||
)
|
||||
messages.append(tool_message)
|
||||
risk_metadata = tool_message.get("_tool_output_risk")
|
||||
if not _flush_session_db_after_tool_progress(
|
||||
|
||||
106
tests/run_agent/test_sequential_tool_timeout.py
Normal file
106
tests/run_agent/test_sequential_tool_timeout.py
Normal file
@@ -0,0 +1,106 @@
|
||||
"""Sequential tool calls recover when one dispatch never returns."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from agent.tool_executor import execute_tool_calls_sequential
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _make_agent(tmp_path: Path) -> AIAgent:
|
||||
with (
|
||||
patch(
|
||||
"run_agent.get_tool_definitions",
|
||||
return_value=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_extract",
|
||||
"description": "test tool",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
),
|
||||
patch("run_agent.check_toolset_requirements", return_value={}),
|
||||
patch("run_agent.OpenAI"),
|
||||
patch("run_agent._hermes_home", tmp_path),
|
||||
patch("agent.model_metadata.fetch_model_metadata", return_value={}),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent._flush_messages_to_session_db = MagicMock(return_value=True)
|
||||
agent._append_guardrail_observation = MagicMock(
|
||||
side_effect=lambda _name, _args, result, **_kwargs: result
|
||||
)
|
||||
agent._record_file_mutation_result = MagicMock()
|
||||
agent._subdirectory_hints.check_tool_call = MagicMock(return_value="")
|
||||
agent._tool_result_content_for_active_model = MagicMock(
|
||||
side_effect=lambda _name, result: result
|
||||
)
|
||||
return agent
|
||||
|
||||
|
||||
def _tool_call(call_id: str):
|
||||
return SimpleNamespace(
|
||||
id=call_id,
|
||||
type="function",
|
||||
function=SimpleNamespace(name="web_extract", arguments="{}"),
|
||||
)
|
||||
|
||||
|
||||
def test_sequential_tool_timeout_emits_result_and_continues(tmp_path, monkeypatch):
|
||||
agent = _make_agent(tmp_path)
|
||||
first_started = threading.Event()
|
||||
release_first = threading.Event()
|
||||
dispatched: list[str] = []
|
||||
terminal_events: list[dict] = []
|
||||
|
||||
def _dispatch(_name, _args, _task_id, *, tool_call_id, **_kwargs):
|
||||
dispatched.append(tool_call_id)
|
||||
if tool_call_id == "hung":
|
||||
first_started.set()
|
||||
release_first.wait()
|
||||
return "late result"
|
||||
return "second result"
|
||||
|
||||
def _capture_terminal_event(*_args, **kwargs):
|
||||
terminal_events.append(kwargs)
|
||||
|
||||
calls = [_tool_call("hung"), _tool_call("next")]
|
||||
assistant = SimpleNamespace(tool_calls=calls)
|
||||
messages: list[dict] = []
|
||||
monkeypatch.setenv("HERMES_CONCURRENT_TOOL_TIMEOUT_S", "0.05")
|
||||
|
||||
started = time.monotonic()
|
||||
try:
|
||||
with (
|
||||
patch("run_agent.handle_function_call", side_effect=_dispatch),
|
||||
patch(
|
||||
"agent.tool_executor._emit_terminal_post_tool_call",
|
||||
side_effect=_capture_terminal_event,
|
||||
),
|
||||
):
|
||||
execute_tool_calls_sequential(agent, assistant, messages, "task")
|
||||
finally:
|
||||
release_first.set()
|
||||
|
||||
assert first_started.is_set()
|
||||
assert time.monotonic() - started < 1.0
|
||||
assert dispatched == ["hung", "next"]
|
||||
assert [message["tool_call_id"] for message in messages] == ["hung", "next"]
|
||||
assert "timed out after 0.1s" in messages[0]["content"]
|
||||
assert messages[0]["effect_disposition"] == "unknown"
|
||||
assert messages[1]["content"] == "second result"
|
||||
timeout_events = [event for event in terminal_events if event.get("error_type") == "tool_timeout"]
|
||||
assert len(timeout_events) == 1
|
||||
assert timeout_events[0]["status"] == "timeout"
|
||||
agent._flush_messages_to_session_db.assert_called()
|
||||
Reference in New Issue
Block a user