Skip to content

Commit 006cfc8

Browse files
authored
Update test_tool_rounds.py
1 parent f56ad89 commit 006cfc8

1 file changed

Lines changed: 132 additions & 0 deletions

File tree

‎tests/agent/test_tool_rounds.py‎

Lines changed: 132 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -239,3 +239,135 @@ def test_salvage_cuts_open_round_variants(self):
239239
],
240240
)
241241
self.assertEqual([m.role for m in loop._salvage_messages()], ["user"])
242+
243+
244+
class TestParallelReadonly(unittest.TestCase):
245+
"""When every call in a round is readonly (Read, Glob, Grep, Skill),
246+
the runner dispatches them concurrently via a thread pool: peak
247+
concurrency equals the call count, wall time is roughly the slowest
248+
tool (not the serial sum), and results are still delivered in
249+
original call order."""
250+
251+
def test_all_readonly_round_runs_in_parallel(self):
252+
"""Three readonly tools (Read, Grep, Glob) with staggered
253+
durations run concurrently: peak concurrency is 3, wall time
254+
is ~max(0.5, 0.1, 0.2)=0.5s (not the 0.8s serial sum), and
255+
results are delivered in original call order."""
256+
session = StaggeredSession({"Read": 0.5, "Grep": 0.1, "Glob": 0.2})
257+
session.tools_enabled = False
258+
session.client.script = [
259+
(
260+
"",
261+
[
262+
ToolCall(id="1", name="Read", arguments='{"file_path": "/tmp/a.py"}'),
263+
ToolCall(id="2", name="Grep", arguments='{"regex": "x", "path": "/tmp"}'),
264+
ToolCall(id="3", name="Glob", arguments='{"pattern": "*.py"}'),
265+
],
266+
),
267+
"done",
268+
]
269+
loop = AgentLoop(session, messages=[Message(role="user", content="go")])
270+
start = time.monotonic()
271+
result = loop.run()
272+
elapsed = time.monotonic() - start
273+
self.assertEqual(result, "done")
274+
self.assertEqual(session.max_active, 3)
275+
self.assertLess(elapsed, 0.8)
276+
self.assertGreaterEqual(elapsed, 0.45)
277+
self.assertEqual(
278+
[m.tool_call_id for m in loop.messages if m.role == "tool"],
279+
["1", "2", "3"],
280+
)
281+
by_id = {m.tool_call_id: m.text() for m in loop.messages if m.role == "tool"}
282+
self.assertEqual(by_id["1"], "result of Read")
283+
self.assertEqual(by_id["2"], "result of Grep")
284+
self.assertEqual(by_id["3"], "result of Glob")
285+
286+
def test_mixed_round_stays_sequential(self):
287+
"""A round with any non-readonly tool (Bash) falls back to
288+
sequential dispatch: peak concurrency is 1, wall time is the
289+
serial sum."""
290+
session = StaggeredSession({"Read": 0.3, "Bash": 0.2, "Grep": 0.1})
291+
session.tools_enabled = False
292+
session.client.script = [
293+
(
294+
"",
295+
[
296+
ToolCall(id="1", name="Read", arguments='{"file_path": "/tmp/a.py"}'),
297+
ToolCall(id="2", name="Bash", arguments='{"command": "echo hi"}'),
298+
ToolCall(id="3", name="Grep", arguments='{"regex": "x", "path": "/tmp"}'),
299+
],
300+
),
301+
"done",
302+
]
303+
loop = AgentLoop(session, messages=[Message(role="user", content="go")])
304+
start = time.monotonic()
305+
result = loop.run()
306+
elapsed = time.monotonic() - start
307+
self.assertEqual(result, "done")
308+
self.assertEqual(session.max_active, 1)
309+
self.assertGreaterEqual(elapsed, 0.55)
310+
311+
def test_single_readonly_tool_runs(self):
312+
"""A single readonly tool still works (no parallelism needed,
313+
but the code path must handle len(calls)==1)."""
314+
session = StaggeredSession({"Read": 0.1})
315+
session.tools_enabled = False
316+
session.client.script = [
317+
(
318+
"",
319+
[ToolCall(id="1", name="Read", arguments='{"file_path": "/tmp/a.py"}')],
320+
),
321+
"done",
322+
]
323+
loop = AgentLoop(session, messages=[Message(role="user", content="go")])
324+
result = loop.run()
325+
self.assertEqual(result, "done")
326+
self.assertEqual(session.max_active, 1)
327+
by_id = {m.tool_call_id: m.text() for m in loop.messages if m.role == "tool"}
328+
self.assertEqual(by_id["1"], "result of Read")
329+
330+
def test_readonly_round_cancel_before_start(self):
331+
"""Ctrl-C before a readonly round starts must skip all tools."""
332+
session = StaggeredSession({"Read": 0.1, "Grep": 0.1})
333+
session.tools_enabled = False
334+
session.cancel()
335+
loop = AgentLoop(session, messages=[Message(role="user", content="go")])
336+
loop.pending = [
337+
ToolCall(id="1", name="Read", arguments='{"file_path": "/tmp/a.py"}'),
338+
ToolCall(id="2", name="Grep", arguments='{"regex": "x", "path": "/tmp"}'),
339+
]
340+
loop._run_tool_round()
341+
self.assertEqual(session.max_active, 0)
342+
self.assertFalse(any(m.role == "tool" for m in loop.messages))
343+
344+
def test_large_readonly_round_caps_concurrency(self):
345+
"""A readonly round larger than MAX_PARALLEL_READONLY must not
346+
spawn one thread per call: peak concurrency is capped, yet
347+
every call still runs and results are delivered in order."""
348+
from python_agent_harness.tool_runner import MAX_PARALLEL_READONLY
349+
350+
n = MAX_PARALLEL_READONLY + 4
351+
session = ParallelToolSession(duration=0.1)
352+
session.tools_enabled = False
353+
session.client.script = [
354+
(
355+
"",
356+
[
357+
ToolCall(id=str(i), name="Read", arguments='{"file_path": "/tmp/a.py"}')
358+
for i in range(n)
359+
],
360+
),
361+
"done",
362+
]
363+
loop = AgentLoop(session, messages=[Message(role="user", content="go")])
364+
result = loop.run()
365+
self.assertEqual(result, "done")
366+
# every call ran, but never more than the cap at once
367+
self.assertEqual(session.executed_count, n)
368+
self.assertEqual(session.max_active, MAX_PARALLEL_READONLY)
369+
# results delivered in original call order
370+
self.assertEqual(
371+
[m.tool_call_id for m in loop.messages if m.role == "tool"],
372+
[str(i) for i in range(n)],
373+
)

0 commit comments

Comments
 (0)