Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions agent/agent_runtime_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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

Expand Down
85 changes: 85 additions & 0 deletions tests/agent/test_tool_pairing_scaling.py
Original file line number Diff line number Diff line change
@@ -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"
))