fix(agent): bound sequential tool calls

This commit is contained in:
fangliquanflq
2026-08-13 03:45:55 +08:00
committed by kshitij
parent 4fe5090964
commit ededa8c4f1
2 changed files with 235 additions and 1 deletions

View File

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

View 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()