fix(delegation): restore parent cancellation for rejected async units

This commit is contained in:
illidan
2026-09-07 12:51:13 +08:00
committed by Teknium
parent 2178b3ebe8
commit 946111cb11
2 changed files with 326 additions and 6 deletions
+22 -6
View File
@@ -14,7 +14,7 @@ from dataclasses import dataclass, replace
from typing import Any, Dict, List, Optional
from tools.async_delegation import _new_delegation_id, record_unit_child
from tools.delegate_tool_child_run import _detach_child, _fabricated_entry, _signal_child_stop
from tools.delegate_tool_child_run import _attach_child, _detach_child, _fabricated_entry, _signal_child_stop
from tools.delegate_tool_progress import (
SUBAGENT_FAILURE_STATUSES, _clean_error_text, _print_completion_line, _quiet, format_batch_tag,
)
@@ -374,6 +374,21 @@ def _dispatch_unit(unit: _Batch, unit_id: Optional[str], slot_key: Optional[str]
progress_fn=lambda: _batch_progress_token(child_agents), **routing,
)
def _restore_parent_cancellation(unit: _Batch) -> None:
# Rejected children remain owned by the parent. Attach before replaying a
# cancellation that may have arrived while async admission had them detached.
parent = unit.parent_agent
for _, _, child in unit.children:
_attach_child(parent, child)
if getattr(parent, "_interrupt_requested", False) is True:
hard_stop = getattr(parent, "_hard_interrupt_requested", None)
for _, _, child in unit.children:
if hard_stop is not None and hard_stop.is_set():
_signal_child_stop(child, getattr(parent, "_interrupt_message", None))
else:
with _quiet("Failed to propagate interrupt to fallback child: %s"):
child.interrupt(getattr(parent, "_interrupt_message", None))
def _dispatch_background(batch: _Batch) -> str:
"""Dispatch the call as independent async units (see ``_units_of``) and return the tool result JSON. Every unit
of one call shares ONE pool slot (``slot_key``), so grouping never changes capacity accounting. Falls back to
@@ -387,10 +402,6 @@ def _dispatch_background(batch: _Batch) -> str:
parent_agent = batch.parent_agent
session_key, origin_ui_session_id = _resolve_async_session_key(parent_agent, batch.origin_ui_session_id)
# The children's lifecycle is owned by the async registry now: drop them from the parent's
# interrupt-propagation list (_build_child_agent attached them, which is correct for sync runs).
for (_, _, c) in batch.children:
_detach_child(parent_agent, c)
routing = dict(
session_key=session_key, origin_ui_session_id=origin_ui_session_id, origin_session_id=wake_sid,
parent_session_id=getattr(parent_agent, "session_id", None), max_async_children=_get_max_async_children(),
@@ -405,11 +416,16 @@ def _dispatch_background(batch: _Batch) -> str:
# cache/delegation/live/<id>/; several units suffix it (-1, -2, ...) and the call keeps the bare id.
unit_id = batch.live_deleg_id if len(units) == 1 else (f"{batch.live_deleg_id}-{k + 1}" if batch.live_deleg_id else None)
unit.unit_id = unit_id = unit_id or _new_delegation_id() # fixed before the runner can start
# The worker can start before admission returns. Detach only this unit:
# unsubmitted units must still receive parent stops while a fallback runs.
for _, _, child in unit.children:
_detach_child(parent_agent, child)
dispatch = _dispatch_unit(unit, unit_id, slot_key, routing)
if dispatch.get("status") == "dispatched":
slot_key = slot_key or dispatch["delegation_id"]
dispatched.append((unit, dispatch["delegation_id"]))
continue
_restore_parent_cancellation(unit)
if not dispatched:
logger.info(
"delegate_task: async pool at capacity (%s); running the whole batch synchronously instead.",
@@ -419,7 +435,7 @@ def _dispatch_background(batch: _Batch) -> str:
# Later units of an admitted call share its slot and cannot be capacity-rejected; a scheduler failure runs
# the unit inline so no task is silently dropped.
logger.warning("delegate_task: unit %d/%d not accepted (%s); running it inline.", k + 1, len(units), dispatch.get("error"))
inline_results.extend(_execute_and_aggregate(unit, honor_parent_interrupt=False)["results"])
inline_results.extend(_execute_and_aggregate(unit)["results"])
payload = _dispatched_payload(batch, dispatched)
if inline_results:
payload["inline_results"] = inline_results