From 6879a69a7195c26162c63ba9474a07638e35e839 Mon Sep 17 00:00:00 2001 From: illidan Date: Tue, 8 Sep 2026 15:41:11 +0800 Subject: [PATCH] perf(agent): avoid quadratic work in tool history pairing --- agent/agent_runtime_helpers.py | 9 +-- tests/agent/test_tool_pairing_scaling.py | 85 ++++++++++++++++++++++++ 2 files changed, 90 insertions(+), 4 deletions(-) create mode 100644 tests/agent/test_tool_pairing_scaling.py diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index f097dee742e2..273aed16e4a7 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -499,7 +499,8 @@ def _prune_unanswered_tool_calls(messages: List[Dict]) -> Tuple[List[Dict], int] pruned.append(msg) continue answered: set = set() - for follower in messages[i + 1:]: + for follower_idx in range(i + 1, len(messages)): + follower = messages[follower_idx] if not (isinstance(follower, dict) and follower.get("role") == "tool"): break tid = (follower.get("tool_call_id") or "").strip() @@ -2512,10 +2513,10 @@ def _classify_tool_call_orphans(messages: List[Dict[str, Any]]): ] result_call_ids: set[str] = set().union(*(v for _, v in result_entries)) orphaned_results = [msg for msg, v in result_entries if v and not (v & surviving_call_ids)] - orphaned_ids = {id(msg) for msg in orphaned_results} - surviving_result_variants = [v for msg, v in result_entries if v and id(msg) not in orphaned_ids] + # Orphan result variants are disjoint from every declared call, so they + # cannot contribute a match. Reuse the union instead of scanning each result. missing_tool_calls = [ - tc for tc, v in assistant_call_variants if not any(v & rv for rv in surviving_result_variants) + tc for tc, v in assistant_call_variants if not (v & result_call_ids) ] return surviving_call_ids, result_call_ids, orphaned_results, missing_tool_calls diff --git a/tests/agent/test_tool_pairing_scaling.py b/tests/agent/test_tool_pairing_scaling.py new file mode 100644 index 000000000000..dca6e16137e4 --- /dev/null +++ b/tests/agent/test_tool_pairing_scaling.py @@ -0,0 +1,85 @@ +"""Bound pairing work by history size without wall-clock timing assertions.""" + +import pytest + +from agent import agent_runtime_helpers as helpers + + +def _exchange(index): + call = { + "id": f"item_{index}", + "call_id": f"call_{index}", + "response_item_id": f"fc_{index}", + "type": "function", + "function": {"name": "read_file", "arguments": "{}"}, + } + return [ + {"role": "assistant", "content": "", "tool_calls": [call]}, + {"role": "tool", "tool_call_id": f"call_{index}|fc_{index}", "content": "ok"}, + ] + + +@pytest.mark.parametrize("pair_count", [64, 256]) +def test_positional_pairing_bounds_history_reads(pair_count): + class CountedHistory(list): + items_read = 0 + + def __iter__(self): + for item in super().__iter__(): + self.items_read += 1 + yield item + + def __getitem__(self, key): + value = super().__getitem__(key) + # A slice copies every selected reference before the caller can break. + self.items_read += len(value) if isinstance(key, slice) else 1 + return value + + paired = [msg for i in range(pair_count) for msg in _exchange(i)] + displaced_call, displaced_result = _exchange("displaced") + boundary = {"role": "user", "content": "next turn"} + messages = CountedHistory(paired + [displaced_call, boundary, displaced_result]) + + repaired, repairs = helpers._prune_unanswered_tool_calls(messages) + + assert messages.items_read <= 8 * len(messages) + assert repairs == 1 + expected = paired + [boundary, displaced_result] + assert repaired == expected + assert all(actual is original for actual, original in zip(repaired, expected)) + + +@pytest.mark.parametrize("pair_count", [64, 256]) +def test_global_pairing_bounds_id_comparisons(pair_count, monkeypatch): + intersections = 0 + original_variants = helpers.tool_call_id_variants + + class CountedVariants(frozenset): + def __and__(self, other): + nonlocal intersections + intersections += 1 + return super().__and__(other) + + monkeypatch.setattr( + helpers, "tool_call_id_variants", + lambda call: CountedVariants(original_variants(call)), + ) + messages = [msg for i in range(pair_count) for msg in _exchange(i)] + unanswered, _ = _exchange("missing") + missing_call = unanswered["tool_calls"][0] + orphan = {"role": "tool", "tool_call_id": "orphan", "content": "unmatched"} + messages.extend([unanswered, orphan]) + + call_ids, result_ids, orphaned, missing = helpers._classify_tool_call_orphans(messages) + + assert intersections <= 4 * (pair_count + 1) + assert orphaned == [orphan] and orphaned[0] is orphan + assert missing == [missing_call] and missing[0] is missing_call + assert call_ids == set().union(*( + original_variants(call) + for message in messages for call in message.get("tool_calls", []) + )) + assert result_ids == set().union(*( + helpers.tool_result_id_variants(message["tool_call_id"]) + for message in messages if message["role"] == "tool" + ))