From ededa8c4f1e211ca3c750c82ad77530dacc0639e Mon Sep 17 00:00:00 2001 From: fangliquanflq Date: Thu, 13 Aug 2026 03:45:55 +0800 Subject: [PATCH] fix(agent): bound sequential tool calls --- agent/tool_executor.py | 130 +++++++++++++++++- .../run_agent/test_sequential_tool_timeout.py | 106 ++++++++++++++ 2 files changed, 235 insertions(+), 1 deletion(-) create mode 100644 tests/run_agent/test_sequential_tool_timeout.py diff --git a/agent/tool_executor.py b/agent/tool_executor.py index cb28c81ea6..a726d5761b 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -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( diff --git a/tests/run_agent/test_sequential_tool_timeout.py b/tests/run_agent/test_sequential_tool_timeout.py new file mode 100644 index 0000000000..96e0a2e87b --- /dev/null +++ b/tests/run_agent/test_sequential_tool_timeout.py @@ -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()