perf(agent): segment mixed tool batches to recover lost concurrency (#64460)

A model response containing several parallel-safe reads plus one unsafe
tool used to lose ALL concurrency: _should_parallelize_tool_batch was
all-or-nothing, so a single barrier call (terminal, clarify, unknown
tool, malformed args) forced the entire batch onto the sequential path.

_plan_tool_batch_segments now splits the batch into ordered segments:
maximal contiguous runs of parallel-safe calls execute on the existing
concurrent path, barrier calls on the sequential path, strictly in the
model's emission order. Invariants preserved:

- one tool result per call, appended in emission order (segments are
  contiguous, so no result reordering across a barrier)
- side-effect boundaries: no call starts before an earlier barrier ends
- overlapping file targets split into separate ordered parallel runs
- turn-end budget enforcement + /steer injection run exactly once per
  batch (segment executors run with finalize=False; the segmented
  dispatcher owns the whole-turn finalize)
- interrupt during segment k drains segments k+1..n with cancelled
  results, keeping one result per tool_call_id

Homogeneous batches keep their original single-path dispatch (zero
behavior delta); _should_parallelize_tool_batch remains as a thin view
over the planner for existing callers and tests.
This commit is contained in:
Teknium
2026-07-14 11:53:05 -07:00
committed by GitHub
parent 7bb409b2d0
commit 271a9d8ec6
4 changed files with 584 additions and 32 deletions

View File

@@ -102,50 +102,118 @@ def _is_mcp_tool_parallel_safe(tool_name: str) -> bool:
return False
def _should_parallelize_tool_batch(tool_calls) -> bool:
"""Return True when a tool-call batch is safe to run concurrently."""
if len(tool_calls) <= 1:
return False
def _plan_tool_batch_segments(tool_calls) -> List[tuple]:
"""Split a tool-call batch into ordered ``(kind, calls)`` segments.
tool_names = [tc.function.name for tc in tool_calls]
if any(name in _NEVER_PARALLEL_TOOLS for name in tool_names):
return False
``kind`` is ``"parallel"`` (a maximal contiguous run of parallel-safe
calls) or ``"sequential"`` (one or more barrier calls that must run
in-order on the sequential path). Segments preserve the model's
original call order exactly — a later call never crosses an earlier
barrier — so tool-result ordering and side-effect boundaries are
identical to fully-sequential execution. The per-call safety rules
are the same ones the old all-or-nothing gate applied to the whole
batch:
* ``_NEVER_PARALLEL_TOOLS`` (interactive tools) → barrier.
* Unparseable / non-dict arguments → barrier.
* Path-scoped tools (``read_file``/``write_file``/``patch``) join a
parallel run only when their target path does not overlap another
path already reserved in the same run; an overlap closes the run so
the conflicting call starts a NEW run after the first completes.
* Anything not in ``_PARALLEL_SAFE_TOOLS`` and not an opted-in MCP
tool → barrier.
Parallel runs shorter than two calls are demoted to sequential (no
concurrency win, and the sequential executor owns the richer inline
dispatch), and adjacent sequential segments are merged.
"""
segments: list[list] = [] # [kind, calls] pairs, normalized to tuples on return
current: list = []
reserved_paths: list[Path] = []
def _close_parallel() -> None:
nonlocal current, reserved_paths
if current:
segments.append(["parallel", current])
current = []
reserved_paths = []
def _add_sequential(tc) -> None:
_close_parallel()
if segments and segments[-1][0] == "sequential":
segments[-1][1].append(tc)
else:
segments.append(["sequential", [tc]])
for tool_call in tool_calls:
tool_name = tool_call.function.name
if tool_name in _NEVER_PARALLEL_TOOLS:
_add_sequential(tool_call)
continue
try:
function_args = json.loads(tool_call.function.arguments)
except Exception:
logging.debug(
"Could not parse args for %s — defaulting to sequential; raw=%s",
"Could not parse args for %s — treating as sequential barrier; raw=%s",
tool_name,
tool_call.function.arguments[:200],
)
return False
_add_sequential(tool_call)
continue
if not isinstance(function_args, dict):
logging.debug(
"Non-dict args for %s (%s) — defaulting to sequential",
"Non-dict args for %s (%s) — treating as sequential barrier",
tool_name,
type(function_args).__name__,
)
return False
_add_sequential(tool_call)
continue
if tool_name in _PATH_SCOPED_TOOLS:
scoped_path = _extract_parallel_scope_path(tool_name, function_args)
if scoped_path is None:
return False
_add_sequential(tool_call)
continue
if any(_paths_overlap(scoped_path, existing) for existing in reserved_paths):
return False
# Same-subtree conflict inside this run: close it so this
# call starts a fresh run AFTER the conflicting one lands.
_close_parallel()
reserved_paths.append(scoped_path)
current.append(tool_call)
continue
if tool_name not in _PARALLEL_SAFE_TOOLS:
# Check if it's an MCP tool from a server that opted into parallel calls.
if not _is_mcp_tool_parallel_safe(tool_name):
return False
if tool_name in _PARALLEL_SAFE_TOOLS or _is_mcp_tool_parallel_safe(tool_name):
current.append(tool_call)
continue
return True
_add_sequential(tool_call)
_close_parallel()
normalized: list[list] = []
for kind, calls in segments:
if kind == "parallel" and len(calls) < 2:
kind = "sequential"
if normalized and normalized[-1][0] == "sequential" and kind == "sequential":
normalized[-1][1].extend(calls)
else:
normalized.append([kind, calls])
return [(kind, calls) for kind, calls in normalized]
def _should_parallelize_tool_batch(tool_calls) -> bool:
"""Return True when the WHOLE tool-call batch is safe to run concurrently.
Thin view over ``_plan_tool_batch_segments`` kept for callers/tests that
only care about the homogeneous case: True iff the planner produces a
single all-parallel segment.
"""
if len(tool_calls) <= 1:
return False
segments = _plan_tool_batch_segments(tool_calls)
return len(segments) == 1 and segments[0][0] == "parallel"
def _extract_parallel_scope_path(tool_name: str, function_args: dict) -> Optional[Path]:
@@ -542,6 +610,7 @@ __all__ = [
"_DESTRUCTIVE_PATTERNS",
"_REDIRECT_OVERWRITE",
"_is_destructive_command",
"_plan_tool_batch_segments",
"_should_parallelize_tool_batch",
"_extract_parallel_scope_path",
"_paths_overlap",