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:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user