refactor(agent/tool_executor): fold observe into _commit_tool_result; _ToolOutcome carries its ref; merge sequential abandon branches
This commit is contained in:
@@ -825,33 +825,28 @@ def _run_sequential_tool_execution_middleware(
|
||||
concurrent.futures.wait([future], timeout=3.0)
|
||||
if future.done() and not future.cancelled():
|
||||
return future.result()
|
||||
abandoned = True
|
||||
future.cancel()
|
||||
interrupt_reason = getattr(agent, "_tool_interrupt_reason", None) or "interrupt requested"
|
||||
message = f"[Tool execution cancelled — {function_name} was abandoned: {interrupt_reason}]"
|
||||
logger.info(
|
||||
"sequential tool %s abandoned due to %s (%.1fs elapsed)",
|
||||
function_name, interrupt_reason, time.monotonic() - started,
|
||||
)
|
||||
return _abandoned_sequential_result(
|
||||
agent, ref, message, _ToolCancelledResult,
|
||||
duration_ms=int((time.monotonic() - started) * 1000),
|
||||
status="cancelled",
|
||||
error_type="tool_interrupted",
|
||||
error_message=f"Tool execution cancelled: {interrupt_reason}",
|
||||
result_cls, outcome = _ToolCancelledResult, dict(
|
||||
duration_ms=int((time.monotonic() - started) * 1000), status="cancelled",
|
||||
error_type="tool_interrupted", error_message=f"Tool execution cancelled: {interrupt_reason}",
|
||||
)
|
||||
else:
|
||||
assert timeout_s is not None # only reachable when a deadline exists
|
||||
message = f"Error executing tool '{function_name}': timed out after {timeout_s:.1f}s"
|
||||
logger.warning("sequential tool %s timed out after %.1fs", function_name, timeout_s)
|
||||
result_cls, outcome = _ToolTimeoutResult, dict(
|
||||
duration_ms=int(timeout_s * 1000), status="timeout", error_type="tool_timeout", error_message=message,
|
||||
)
|
||||
|
||||
# Only reachable when a deadline exists (interrupted returns above).
|
||||
assert timeout_s is not None
|
||||
abandoned = True
|
||||
message = f"Error executing tool '{function_name}': timed out after {timeout_s:.1f}s"
|
||||
logger.warning("sequential tool %s timed out after %.1fs", function_name, timeout_s)
|
||||
future.cancel()
|
||||
_interrupt_worker_tids(agent, worker_tid)
|
||||
return _abandoned_sequential_result(
|
||||
agent, ref, message, _ToolTimeoutResult,
|
||||
duration_ms=int(timeout_s * 1000), status="timeout", error_type="tool_timeout", error_message=message,
|
||||
)
|
||||
if state == "timeout":
|
||||
_interrupt_worker_tids(agent, worker_tid)
|
||||
return _abandoned_sequential_result(agent, ref, message, result_cls, **outcome)
|
||||
finally:
|
||||
# Never join a wedged worker (daemon pool also keeps it out of the atexit join).
|
||||
executor.shutdown(wait=not abandoned, cancel_futures=abandoned)
|
||||
@@ -935,12 +930,39 @@ def _commit_tool_result(
|
||||
is_error: bool,
|
||||
blocked: bool,
|
||||
effect_disposition,
|
||||
observed: bool = False,
|
||||
error_preview: Callable[[Any], Any] = lambda result: result,
|
||||
success_log_chars: Optional[int] = None,
|
||||
verbose_text: Callable[[Any], Any] = lambda result: result,
|
||||
):
|
||||
"""Mark the tool done; persist/spill, hint, wrap and append its result; flush the session
|
||||
DB; then project ``tool.completed``. Returns ``(persisted_result, display_result,
|
||||
risk_metadata)`` (``display_result`` = pre-persist content for UI previews) or ``None``
|
||||
when the flush failed (the caller must stop the batch)."""
|
||||
"""Observe (``observed`` results only) and log the outcome; mark the tool done; persist/
|
||||
spill, hint, wrap and append the result; flush the session DB; project ``tool.completed``.
|
||||
|
||||
Blocked calls never ran, so they are neither guardrail-observed nor fed to the file-
|
||||
mutation verifier; ``success_log_chars`` (sequential path) also logs the completion line.
|
||||
Returns ``(persisted_result, display_result, risk_metadata)`` (``display_result`` =
|
||||
pre-persist content for UI previews) or ``None`` when the flush failed (stop the batch).
|
||||
"""
|
||||
function_name, function_args, tool_call_id, effective_task_id = ref.name, ref.args, ref.call_id, ref.task_id
|
||||
if observed:
|
||||
if not blocked:
|
||||
function_result = agent._append_guardrail_observation(
|
||||
function_name, function_args, function_result, failed=is_error, tool_call_id=tool_call_id,
|
||||
)
|
||||
if is_error:
|
||||
logger.warning("Tool %s returned error (%.2fs): %s", function_name, tool_duration, error_preview(function_result))
|
||||
elif success_log_chars is not None:
|
||||
logger.info("tool %s completed (%.2fs, %d chars)", function_name, tool_duration, success_log_chars)
|
||||
if not blocked:
|
||||
try:
|
||||
agent._record_file_mutation_result(function_name, function_args, function_result, is_error)
|
||||
except Exception as _ver_err:
|
||||
logging.debug("file-mutation verifier record failed: %s", _ver_err)
|
||||
if agent.verbose_logging:
|
||||
logging.debug("Tool %s completed in %.2fs", function_name, tool_duration)
|
||||
_log_result = verbose_text(function_result)
|
||||
logging.debug("Tool result (%d chars): %s", len(_log_result), _log_result)
|
||||
|
||||
agent._current_tool = None
|
||||
_status_suffix = " (error)" if is_error else ""
|
||||
agent._touch_activity(f"tool completed: {function_name} ({tool_duration:.1f}s){_status_suffix}")
|
||||
@@ -982,52 +1004,6 @@ def _commit_tool_result(
|
||||
return persisted_result, function_result, tool_message.get("_tool_output_risk")
|
||||
|
||||
|
||||
def _observe_and_commit_tool_result(
|
||||
agent,
|
||||
messages: list,
|
||||
ref: _ToolCallRef,
|
||||
function_result,
|
||||
*,
|
||||
budget: BudgetConfig,
|
||||
tool_duration: float,
|
||||
is_error: bool,
|
||||
blocked: bool,
|
||||
effect_disposition,
|
||||
error_preview: Callable[[Any], Any],
|
||||
success_log_chars: Optional[int] = None,
|
||||
verbose_text: Callable[[Any], Any] = lambda result: result,
|
||||
):
|
||||
"""Guardrail-observe a result that actually ran, log its outcome, feed the turn-end
|
||||
file-mutation verifier, then ``_commit_tool_result`` it. Blocked calls never ran, so
|
||||
they count as neither failure nor success and are not observed; ``success_log_chars``
|
||||
(sequential path) also logs the completion line."""
|
||||
function_name, function_args, tool_call_id = ref.name, ref.args, ref.call_id
|
||||
if not blocked:
|
||||
function_result = agent._append_guardrail_observation(
|
||||
function_name, function_args, function_result, failed=is_error, tool_call_id=tool_call_id,
|
||||
)
|
||||
if is_error:
|
||||
logger.warning("Tool %s returned error (%.2fs): %s", function_name, tool_duration, error_preview(function_result))
|
||||
elif success_log_chars is not None:
|
||||
logger.info("tool %s completed (%.2fs, %d chars)", function_name, tool_duration, success_log_chars)
|
||||
if not blocked:
|
||||
try:
|
||||
agent._record_file_mutation_result(function_name, function_args, function_result, is_error)
|
||||
except Exception as _ver_err:
|
||||
logging.debug("file-mutation verifier record failed: %s", _ver_err)
|
||||
|
||||
if agent.verbose_logging:
|
||||
logging.debug("Tool %s completed in %.2fs", function_name, tool_duration)
|
||||
_log_result = verbose_text(function_result)
|
||||
logging.debug("Tool result (%d chars): %s", len(_log_result), _log_result)
|
||||
|
||||
return _commit_tool_result(
|
||||
agent, messages, ref, function_result,
|
||||
budget=budget, tool_duration=tool_duration, is_error=is_error, blocked=blocked,
|
||||
effect_disposition=effect_disposition,
|
||||
)
|
||||
|
||||
|
||||
def _finalize_tool_batch(agent, messages: list, effective_task_id: str, num_tools: int, budget: BudgetConfig) -> None:
|
||||
"""Per-turn aggregate budget enforcement, then /steer injection — in that order, so the
|
||||
steer marker is never truncated/discarded when enforcement replaces a result."""
|
||||
@@ -1057,15 +1033,13 @@ def _print_tool_completed(agent, index: int, tool_duration: float, result) -> No
|
||||
|
||||
@dataclass
|
||||
class _ToolOutcome:
|
||||
"""One finished worker slot of a concurrent batch."""
|
||||
"""One finished worker slot of a concurrent batch (``ref`` holds the final name/args/trace)."""
|
||||
|
||||
name: str
|
||||
args: dict
|
||||
ref: _ToolCallRef
|
||||
result: Any
|
||||
duration: float
|
||||
is_error: bool
|
||||
blocked: bool
|
||||
middleware_trace: list
|
||||
|
||||
|
||||
def _start_order_gate_timeout(batch_timeout: float | None) -> float:
|
||||
@@ -1144,7 +1118,7 @@ class _ConcurrentBatch:
|
||||
self.results: list[Optional[_ToolOutcome]] = [None] * len(parsed_calls)
|
||||
for i, pc in enumerate(parsed_calls):
|
||||
if pc.parse_error is not None:
|
||||
self.results[i] = _ToolOutcome(pc.name, pc.args, pc.parse_error, 0.0, True, True, pc.middleware_trace)
|
||||
self.results[i] = _ToolOutcome(pc.ref(effective_task_id), pc.parse_error, 0.0, True, True)
|
||||
self.gate = _StartOrderGate(_start_order_gate_timeout(timeout_s))
|
||||
self.authorization_gate = _ConcurrentToolAuthorizationGate()
|
||||
self.timed_out_indices: set[int] = set()
|
||||
@@ -1186,7 +1160,7 @@ class _ConcurrentBatch:
|
||||
result = ref.emit_cancelled(agent, start)
|
||||
duration = time.time() - start
|
||||
logger.info("tool %s cancelled (%.2fs)", ref.name, duration)
|
||||
return _ToolOutcome(ref.name, ref.args, result, duration, True, False, ref.trace)
|
||||
return _ToolOutcome(ref, result, duration, True, False)
|
||||
except Exception as tool_error:
|
||||
result = f"Error executing tool '{ref.name}': {tool_error}"
|
||||
logger.error("_invoke_tool raised for %s: %s", ref.name, tool_error, exc_info=True)
|
||||
@@ -1198,7 +1172,7 @@ class _ConcurrentBatch:
|
||||
logger.info("tool %s failed (%.2fs): %s", ref.name, duration, result[:200])
|
||||
else:
|
||||
logger.info("tool %s completed (%.2fs, %d chars)", ref.name, duration, len(result))
|
||||
return _ToolOutcome(ref.name, ref.args, result, duration, is_error, blocked, ref.trace)
|
||||
return _ToolOutcome(ref, result, duration, is_error, blocked)
|
||||
|
||||
def run_worker(self, index: int, start_order: int) -> None:
|
||||
"""Worker function executed in a thread."""
|
||||
@@ -1234,10 +1208,10 @@ class _ConcurrentBatch:
|
||||
"interpreter shutdown while scheduling concurrent tools; skipping %d unsubmitted tool(s)", len(skipped),
|
||||
)
|
||||
for skipped_i in skipped:
|
||||
pc = self.parsed_calls[skipped_i]
|
||||
ref = self.parsed_calls[skipped_i].ref(self.effective_task_id)
|
||||
if self.results[skipped_i] is None:
|
||||
result = f"Error executing tool '{pc.name}': Python interpreter is shutting down; tool was not started"
|
||||
self.results[skipped_i] = _ToolOutcome(pc.name, pc.args, result, 0.0, True, False, pc.middleware_trace)
|
||||
result = f"Error executing tool '{ref.name}': Python interpreter is shutting down; tool was not started"
|
||||
self.results[skipped_i] = _ToolOutcome(ref, result, 0.0, True, False)
|
||||
break
|
||||
futures.append(f)
|
||||
future_to_index[f] = i
|
||||
@@ -1352,29 +1326,24 @@ def _append_batch_results(agent, messages: list, effective_task_id: str, batch:
|
||||
failed flush (the caller must stop the batch)."""
|
||||
for i, pc in enumerate(batch.parsed_calls):
|
||||
r = batch.results[i]
|
||||
ref = pc.ref(effective_task_id)
|
||||
# A worker may finish between the deadline snapshot and this loop;
|
||||
# prefer its real result over a fabricated timeout.
|
||||
if r is None:
|
||||
blocked = False
|
||||
ref, is_error, blocked = pc.ref(effective_task_id), True, False
|
||||
function_result, tool_duration, effect_disposition = _unfinished_tool_result(
|
||||
agent, ref, timed_out=i in batch.timed_out_indices, timeout_s=batch.timeout_s,
|
||||
)
|
||||
committed = _commit_tool_result(
|
||||
agent, messages, ref, function_result,
|
||||
budget=budget, tool_duration=tool_duration, is_error=True, blocked=False,
|
||||
effect_disposition=effect_disposition,
|
||||
)
|
||||
else:
|
||||
ref.name, ref.args, ref.trace, tool_duration, blocked = r.name, r.args, r.middleware_trace, r.duration, r.blocked
|
||||
ref, function_result, tool_duration, is_error, blocked = r.ref, r.result, r.duration, r.is_error, r.blocked
|
||||
effect_disposition = "none" if blocked else None
|
||||
if pc.parse_error is not None:
|
||||
ref.emit_invalid_arguments(agent, r.result)
|
||||
committed = _observe_and_commit_tool_result(
|
||||
agent, messages, ref, r.result,
|
||||
budget=budget, tool_duration=tool_duration, is_error=r.is_error, blocked=blocked,
|
||||
effect_disposition="none" if blocked else None,
|
||||
error_preview=lambda res: _multimodal_text_summary(res)[:200],
|
||||
)
|
||||
committed = _commit_tool_result(
|
||||
agent, messages, ref, function_result,
|
||||
budget=budget, tool_duration=tool_duration, is_error=is_error, blocked=blocked,
|
||||
effect_disposition=effect_disposition, observed=r is not None,
|
||||
error_preview=lambda res: _multimodal_text_summary(res)[:200],
|
||||
)
|
||||
if committed is None:
|
||||
return False
|
||||
_persisted, display_function_result, risk_metadata = committed
|
||||
@@ -1492,19 +1461,13 @@ class _SequentialDispatch:
|
||||
finish_in_finally: bool = True
|
||||
|
||||
|
||||
def _resolve_sequential_dispatch(
|
||||
agent,
|
||||
*,
|
||||
function_name: str,
|
||||
function_args: dict,
|
||||
messages: list,
|
||||
effective_task_id: str,
|
||||
tool_call_id: str,
|
||||
middleware_trace: list,
|
||||
) -> _SequentialDispatch:
|
||||
def _resolve_sequential_dispatch(agent, ref: _ToolCallRef, messages: list) -> _SequentialDispatch:
|
||||
"""Pick the execute callable for one sequential call and start its spinner. Precedence:
|
||||
inline agent-level tools, delegate_task, context-engine tools, memory-provider tools,
|
||||
then the registry."""
|
||||
function_name, function_args, effective_task_id, tool_call_id, middleware_trace = (
|
||||
ref.name, ref.args, ref.task_id, ref.call_id, ref.trace,
|
||||
)
|
||||
if function_name != "delegate_task" and function_name in INLINE_TOOL_EXECUTORS:
|
||||
# Agent-level tools that need live AIAgent state; table shared with invoke_tool.
|
||||
inline_executor = INLINE_TOOL_EXECUTORS[function_name]
|
||||
@@ -1654,10 +1617,10 @@ def _publish_sequential_result(agent, messages: list, ref: _ToolCallRef, managed
|
||||
# observer is suppressed); also stops an abandoned timeout worker reporting late.
|
||||
if not managed.blocked and not _execution_timed_out:
|
||||
ref.emit_post(agent, function_result, duration_ms=int(tool_duration * 1000))
|
||||
committed = _observe_and_commit_tool_result(
|
||||
committed = _commit_tool_result(
|
||||
agent, messages, ref, function_result,
|
||||
budget=budget, tool_duration=tool_duration, is_error=_is_error_result, blocked=managed.blocked,
|
||||
effect_disposition="unknown" if _execution_timed_out else None,
|
||||
effect_disposition="unknown" if _execution_timed_out else None, observed=True,
|
||||
error_preview=lambda res: res[:200] if isinstance(res, str) and not agent.verbose_logging else res,
|
||||
success_log_chars=_result_len,
|
||||
verbose_text=_multimodal_text_summary,
|
||||
@@ -1680,7 +1643,6 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
tool_calls = assistant_message.tool_calls
|
||||
|
||||
for i, tool_call in enumerate(tool_calls, 1):
|
||||
tool_call_id = _pairing_tool_call_id(tool_call)
|
||||
if getattr(agent, "_incremental_persistence_failed", False):
|
||||
return
|
||||
# Check interrupt BEFORE each tool so a "stop" during the previous one skips the rest.
|
||||
@@ -1704,15 +1666,7 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
continue
|
||||
|
||||
tool_start_time = time.time()
|
||||
dispatch = _resolve_sequential_dispatch(
|
||||
agent,
|
||||
function_name=ref.name,
|
||||
function_args=ref.args,
|
||||
messages=messages,
|
||||
effective_task_id=effective_task_id,
|
||||
tool_call_id=tool_call_id,
|
||||
middleware_trace=ref.trace,
|
||||
)
|
||||
dispatch = _resolve_sequential_dispatch(agent, ref, messages)
|
||||
managed, tool_duration = _run_sequential_call(
|
||||
agent, dispatch, ref,
|
||||
scope_block=pc.scope_block,
|
||||
|
||||
Reference in New Issue
Block a user