Skip to content
Merged
37 changes: 36 additions & 1 deletion pyrit/scenario/core/scenario.py
Original file line number Diff line number Diff line change
Expand Up @@ -1687,6 +1687,11 @@ async def _execute_atomic_attacks_parallel_async(
work persists for resume). If more than one in-flight attack ends up failing,
every failure is surfaced: a single failure is re-raised as-is, multiple
failures are wrapped in an ``ExceptionGroup`` so callers see all of them.
Cancellation stops queue admission and cancels and drains all workers before
propagating, including when it originates inside an atomic attack.
Workers already processing cancellation are drained without a second request.
Queue admission observes new supervisor cancellation requests before its
cancellation handler resumes, without treating earlier requests as a new cancellation.
"""
# Type narrowing: initialize_async always sets _max_concurrency to an int. We hold
# the narrowed value in a local so the type checker can verify all uses below.
Expand Down Expand Up @@ -1714,10 +1719,16 @@ async def _execute_atomic_attacks_parallel_async(
queue.put_nowait(atomic_attack)

stop_event = asyncio.Event()
supervisor = asyncio.current_task()
assert supervisor is not None, "Scenario worker pool requires a running task."
initial_cancellations = supervisor.cancelling()
outcomes: list[tuple[AtomicAttack, AttackExecutorResult[AttackResult]] | Exception] = []

async def worker_async() -> None:
while not stop_event.is_set():
if supervisor.cancelling() > initial_cancellations:
stop_event.set()
return
try:
atomic_attack = queue.get_nowait()
except asyncio.QueueEmpty:
Expand All @@ -1735,6 +1746,10 @@ async def worker_async() -> None:
except Exception as exc:
outcomes.append(exc)
stop_event.set()
except BaseException:
# Stop admission before a ready sibling can take another queued attack.
stop_event.set()
raise
finally:
self._active_atomic_groups.pop(atomic_group_id, None)
pbar.update(1)
Expand All @@ -1744,8 +1759,28 @@ async def worker_async() -> None:
# without losing parallelism for the common case where remaining_attacks fits in
# the budget.
worker_count = min(max_concurrency, len(remaining_attacks))
workers = [asyncio.create_task(worker_async()) for _ in range(worker_count)]
group = asyncio.gather(*workers)
try:
await asyncio.gather(*(worker_async() for _ in range(worker_count)))
# The supervisor owns cancellation; gather must not forward it ahead of this handler.
await asyncio.shield(group)
Comment thread
romanlutz marked this conversation as resolved.
except BaseException:
# gather does not cancel siblings when a child is cancelled.
stop_event.set()
for worker in workers:
if not worker.done() and not worker.cancelling():
worker.cancel()
drain = asyncio.gather(group, *workers, return_exceptions=True)
caller_cancellation: asyncio.CancelledError | None = None
while not drain.done():
try:
await asyncio.shield(drain)
except asyncio.CancelledError as cancellation:
caller_cancellation = cancellation
drain.result()
if caller_cancellation is not None:
raise caller_cancellation from None
raise
finally:
pbar.close()

Expand Down
35 changes: 35 additions & 0 deletions tests/unit/scenario/core/test_scenario.py
Original file line number Diff line number Diff line change
Expand Up @@ -1981,6 +1981,41 @@ async def side_run_2(*a, **k):
# Sanity check: the failure actually happened.
assert bad_started.is_set()

async def test_child_cancellation_stops_queue_before_ready_sibling_finishes(
self, mock_atomic_attacks, sample_attack_results, mock_objective_target
):
sibling_started = asyncio.Event()
release_sibling = asyncio.Event()

async def cancelled_run_async(**_kwargs):
await sibling_started.wait()
release_sibling.set()
raise asyncio.CancelledError("atomic attack cancelled")

async def sibling_run_async(**_kwargs):
sibling_started.set()
await release_sibling.wait()
return AttackExecutorResult(completed_results=[sample_attack_results[1]], incomplete_objectives=[])

mock_atomic_attacks[0].run_async = AsyncMock(side_effect=cancelled_run_async)
mock_atomic_attacks[1].run_async = AsyncMock(side_effect=sibling_run_async)
mock_atomic_attacks[2].run_async = create_mock_run_async(
[sample_attack_results[2]], atomic_attack=mock_atomic_attacks[2]
)
scenario = ConcreteScenario(
name="Child Cancellation Scenario",
version=1,
atomic_attacks_to_return=mock_atomic_attacks,
)
scenario.set_params_from_args(args={"objective_target": mock_objective_target, "max_concurrency": 2})
await scenario.initialize_async()

with pytest.raises(asyncio.CancelledError, match="atomic attack cancelled"):
await asyncio.wait_for(scenario.run_async(), timeout=5)

mock_atomic_attacks[2].run_async.assert_not_called()
assert not scenario._active_atomic_groups

async def test_multiple_inflight_failures_are_grouped_into_exception_group(
self, mock_atomic_attacks, sample_attack_results, mock_objective_target
):
Expand Down
Loading
Loading