From c9d0accfe3b237688313ca2a6fa5a535fa21c7ac Mon Sep 17 00:00:00 2001 From: biefan <70761325+biefan@users.noreply.github.com> Date: Fri, 25 Sep 2026 19:41:42 +0000 Subject: [PATCH 1/4] FIX Stop scenario workers when an atomic attack is cancelled --- pyrit/scenario/core/scenario.py | 17 ++++++++- tests/unit/scenario/core/test_scenario.py | 35 +++++++++++++++++++ .../core/test_scenario_partial_results.py | 34 +++++++++++++----- 3 files changed, 76 insertions(+), 10 deletions(-) diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 8fb8ce704c..3e1f3f0e7d 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -1752,6 +1752,8 @@ 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. """ # 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. @@ -1800,6 +1802,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) @@ -1809,8 +1815,17 @@ 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)] try: - await asyncio.gather(*(worker_async() for _ in range(worker_count))) + await asyncio.gather(*workers) + except BaseException: + # gather does not cancel siblings when a child is cancelled. + stop_event.set() + for worker in workers: + if not worker.done(): + worker.cancel() + await asyncio.gather(*workers, return_exceptions=True) + raise finally: pbar.close() diff --git a/tests/unit/scenario/core/test_scenario.py b/tests/unit/scenario/core/test_scenario.py index a22d43cca5..cd75eb03f1 100644 --- a/tests/unit/scenario/core/test_scenario.py +++ b/tests/unit/scenario/core/test_scenario.py @@ -1897,6 +1897,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 ): diff --git a/tests/unit/scenario/core/test_scenario_partial_results.py b/tests/unit/scenario/core/test_scenario_partial_results.py index cfb3b1e51a..a754b5c010 100644 --- a/tests/unit/scenario/core/test_scenario_partial_results.py +++ b/tests/unit/scenario/core/test_scenario_partial_results.py @@ -423,7 +423,10 @@ async def mock_run(*args, **kwargs): # All 5 results should be in final scenario result assert len(result.attack_results["resume_attack"]) == 5 - async def test_run_async_cancellation_persists_progress_cleans_workers_and_resumes(self, mock_objective_target): + @pytest.mark.parametrize("cancel_worker", [False, True], ids=["caller-cancelled", "worker-cancelled"]) + async def test_run_async_cancellation_persists_progress_cleans_workers_and_resumes( + self, mock_objective_target, cancel_worker + ): completed_attack = create_mock_atomic_attack("completed_attack", ["obj1"]) in_flight_attack = create_mock_atomic_attack("in_flight_attack", ["obj2"]) queued_attack = create_mock_atomic_attack("queued_attack", ["obj3"]) @@ -455,18 +458,22 @@ async def test_run_async_cancellation_persists_progress_cleans_workers_and_resum in_flight_worker_exited = asyncio.Event() block_until_cancelled = asyncio.Event() persisted_objectives: list[str] = [] + worker_tasks: list[asyncio.Task] = [] async def run_completed_attack(*args, **kwargs): + worker_tasks.append(asyncio.current_task()) save_attack_results_to_memory([completed_result], atomic_attack=completed_attack) persisted_objectives.append(completed_result.objective) completed_persisted.set() try: await block_until_cancelled.wait() finally: + await asyncio.sleep(0) completed_worker_exited.set() async def run_in_flight_attack(*args, **kwargs): if in_flight_attack.run_async.call_count == 1: + worker_tasks.append(asyncio.current_task()) in_flight_started.set() try: await block_until_cancelled.wait() @@ -505,15 +512,24 @@ async def run_queued_attack(*args, **kwargs): scenario_task = asyncio.create_task(scenario.run_async()) await asyncio.wait_for(completed_persisted.wait(), timeout=5.0) await asyncio.wait_for(in_flight_started.wait(), timeout=5.0) - scenario_task.cancel() + task_to_cancel = worker_tasks[1] if cancel_worker else scenario_task + task_to_cancel.cancel() - with pytest.raises(asyncio.CancelledError): - await scenario_task - - assert completed_worker_exited.is_set() - assert in_flight_worker_exited.is_set() - queued_attack.run_async.assert_not_called() - assert persisted_objectives == ["obj1"] + try: + with pytest.raises(asyncio.CancelledError): + await scenario_task + + assert completed_worker_exited.is_set() + assert in_flight_worker_exited.is_set() + assert all(task.done() for task in worker_tasks) + assert not scenario._active_atomic_groups + queued_attack.run_async.assert_not_called() + assert persisted_objectives == ["obj1"] + finally: + for task in worker_tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*worker_tasks, return_exceptions=True) [cancelled_result] = CentralMemory.get_memory_instance().get_scenario_results( scenario_result_ids=[scenario._scenario_result_id] From 70ce82c0fb0b5de3137ae153643c072579ee2291 Mon Sep 17 00:00:00 2001 From: biefan <70761325+biefan@users.noreply.github.com> Date: Fri, 25 Sep 2026 21:43:18 +0000 Subject: [PATCH 2/4] FIX Preserve scenario worker cleanup during repeated cancellation --- pyrit/scenario/core/scenario.py | 11 ++++- .../core/test_scenario_partial_results.py | 45 +++++++++++++++++++ 2 files changed, 55 insertions(+), 1 deletion(-) diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 3e1f3f0e7d..4c88c38098 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -1824,7 +1824,16 @@ async def worker_async() -> None: for worker in workers: if not worker.done(): worker.cancel() - await asyncio.gather(*workers, return_exceptions=True) + drain = asyncio.gather(*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() diff --git a/tests/unit/scenario/core/test_scenario_partial_results.py b/tests/unit/scenario/core/test_scenario_partial_results.py index a754b5c010..3369929bbb 100644 --- a/tests/unit/scenario/core/test_scenario_partial_results.py +++ b/tests/unit/scenario/core/test_scenario_partial_results.py @@ -553,6 +553,51 @@ async def run_queued_attack(*args, **kwargs): assert sorted(resumed_result.get_objectives()) == ["obj1", "obj2", "obj3"] assert all(len(results) == 1 for results in resumed_result.attack_results.values()) + async def test_worker_cancellation_waits_for_cleanup_despite_caller_cancellation(self, mock_objective_target): + cancelled_attack = create_mock_atomic_attack("cancelled", ["obj1"]) + sibling_attack = create_mock_atomic_attack("sibling", ["obj2"]) + sibling_started = asyncio.Event() + cleanup_started = asyncio.Event() + allow_cleanup = asyncio.Event() + cleanup_finished = asyncio.Event() + + async def cancelled_run_async(**_kwargs): + await sibling_started.wait() + raise asyncio.CancelledError("child cancelled") + + async def sibling_run_async(**_kwargs): + sibling_started.set() + try: + await asyncio.Event().wait() + finally: + cleanup_started.set() + await allow_cleanup.wait() + cleanup_finished.set() + + cancelled_attack.run_async = AsyncMock(side_effect=cancelled_run_async) + sibling_attack.run_async = AsyncMock(side_effect=sibling_run_async) + scenario = ConcreteScenario( + name="Cancellation During Cleanup", version=1, atomic_attacks_to_return=[cancelled_attack, sibling_attack] + ) + scenario.set_params_from_args(args={"objective_target": mock_objective_target, "max_concurrency": 2}) + await scenario.initialize_async() + + task = asyncio.create_task(scenario.run_async()) + try: + await asyncio.wait_for(cleanup_started.wait(), timeout=5) + task.cancel("caller cancelled during cleanup") + await asyncio.sleep(0) + allow_cleanup.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=5) + assert cleanup_finished.is_set() + assert not scenario._active_atomic_groups + finally: + allow_cleanup.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + async def test_run_async_cancellation_is_not_masked_by_persistence_failure( self, mock_objective_target: MagicMock ) -> None: From b509c97e20f7c17f2080a4cf0e3ffa2d66a5e38a Mon Sep 17 00:00:00 2001 From: biefan <70761325+biefan@users.noreply.github.com> Date: Sat, 26 Sep 2026 08:57:30 +0000 Subject: [PATCH 3/4] FIX Preserve scenario cleanup when cancellation is already in flight --- pyrit/scenario/core/scenario.py | 9 +- .../core/test_scenario_partial_results.py | 91 ++++++++++++++++++- 2 files changed, 97 insertions(+), 3 deletions(-) diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 4c88c38098..a1940f0a5f 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -1754,6 +1754,7 @@ async def _execute_atomic_attacks_parallel_async( 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. """ # 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. @@ -1816,13 +1817,15 @@ async def worker_async() -> None: # 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(*workers) + # The supervisor owns cancellation; gather must not forward it ahead of this handler. + await asyncio.shield(group) except BaseException: # gather does not cancel siblings when a child is cancelled. stop_event.set() for worker in workers: - if not worker.done(): + if not worker.done() and not worker.cancelling(): worker.cancel() drain = asyncio.gather(*workers, return_exceptions=True) caller_cancellation: asyncio.CancelledError | None = None @@ -1832,6 +1835,8 @@ async def worker_async() -> None: except asyncio.CancelledError as cancellation: caller_cancellation = cancellation drain.result() + # Retrieve an exception that arrived after the initial shield was cancelled. + group.exception() if caller_cancellation is not None: raise caller_cancellation from None raise diff --git a/tests/unit/scenario/core/test_scenario_partial_results.py b/tests/unit/scenario/core/test_scenario_partial_results.py index 3369929bbb..89bc4676c3 100644 --- a/tests/unit/scenario/core/test_scenario_partial_results.py +++ b/tests/unit/scenario/core/test_scenario_partial_results.py @@ -10,6 +10,7 @@ import pytest from pyrit.exceptions import ScenarioPartialFailureException +from pyrit.executor.attack import PromptSendingAttack from pyrit.executor.attack.core import AttackExecutorResult from pyrit.memory import CentralMemory from pyrit.models import ( @@ -23,7 +24,8 @@ ) from pyrit.prompt_target import PromptTarget from pyrit.scenario import DatasetConfiguration, ScenarioResult -from pyrit.scenario.core import AtomicAttack, BaselineAttackPolicy, Scenario, ScenarioTechnique +from pyrit.scenario.core import AtomicAttack, AttackTechnique, BaselineAttackPolicy, Scenario, ScenarioTechnique +from tests.unit.mocks import MockPromptTarget def _mock_scorer_id(name: str = "MockScorer") -> ComponentIdentifier: @@ -598,6 +600,93 @@ async def sibling_run_async(**_kwargs): task.cancel() await asyncio.gather(task, return_exceptions=True) + @pytest.mark.parametrize("cancel_again", [False, True], ids=["single-cancel", "repeated-cancel"]) + async def test_caller_cancellation_drains_real_target_reset_without_recancelling(self, cancel_again): + target = MockPromptTarget() + all_started = asyncio.Event() + fast_worker_finished = asyncio.Event() + cleanup_started = asyncio.Event() + release_cleanup = asyncio.Event() + cleanup_finished = asyncio.Event() + sends: dict[str, asyncio.Task] = {} + conversations: dict[str, str] = {} + + async def send_async(*, normalized_conversation): + piece = normalized_conversation[-1].get_piece() + task = asyncio.current_task() + assert task is not None + sends[piece.converted_value] = task + conversations[piece.conversation_id] = piece.converted_value + if len(sends) == 2: + all_started.set() + await asyncio.Event().wait() + + async def reset_async(*, conversation_id): + if conversations[conversation_id] == "slow": + cleanup_started.set() + await release_cleanup.wait() + cleanup_finished.set() + + atomics = [ + AtomicAttack( + atomic_attack_name=name, + attack_technique=AttackTechnique(attack=PromptSendingAttack(objective_target=target)), + seed_groups=[AttackSeedGroup(seeds=[SeedObjective(value=name)])], + ) + for name in ("slow", "fast", "queued") + ] + scenario = ConcreteScenario(name="Real Target Cleanup", version=1, atomic_attacks_to_return=atomics) + scenario.set_params_from_args(args={"objective_target": target, "max_concurrency": 2, "max_retries": 2}) + await scenario.initialize_async() + fast_run_async = atomics[1].run_async + + async def observe_fast_worker_async(**kwargs): + worker = asyncio.current_task() + assert worker is not None + worker.add_done_callback(lambda _: fast_worker_finished.set()) + return await fast_run_async(**kwargs) + + with ( + patch.object(target, "_send_prompt_to_target_async", new=send_async), + patch.object(target, "reset_conversation_async", new=reset_async), + patch.object(atomics[1], "run_async", new=observe_fast_worker_async), + ): + parent = asyncio.create_task(scenario.run_async()) + try: + await asyncio.wait_for(all_started.wait(), timeout=5) + parent.cancel("stop scenario") + await asyncio.wait_for(cleanup_started.wait(), timeout=5) + await asyncio.wait_for(fast_worker_finished.wait(), timeout=5) + assert sends["slow"].cancelling() == 1 + assert not cleanup_finished.is_set() + assert not parent.done() + [stored] = scenario._memory.get_scenario_results(scenario_result_ids=[scenario._scenario_result_id]) + assert stored.scenario_run_state is ScenarioRunState.IN_PROGRESS + + if cancel_again: + parent.cancel("stop scenario again") + cancellation_delivered = asyncio.Event() + asyncio.get_running_loop().call_soon(cancellation_delivered.set) + await cancellation_delivered.wait() + assert sends["slow"].cancelling() == 1 + assert not parent.done() + + release_cleanup.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(parent, timeout=5) + assert cleanup_finished.is_set() + assert set(sends) == {"slow", "fast"} + assert all(task.done() for task in sends.values()) + assert not scenario._active_atomic_groups + [stored] = scenario._memory.get_scenario_results(scenario_result_ids=[scenario._scenario_result_id]) + assert stored.scenario_run_state is ScenarioRunState.CANCELLED + assert stored.number_tries == 1 + finally: + release_cleanup.set() + if not parent.done(): + parent.cancel() + await asyncio.gather(parent, *sends.values(), return_exceptions=True) + async def test_run_async_cancellation_is_not_masked_by_persistence_failure( self, mock_objective_target: MagicMock ) -> None: From c79c8226c478479f333bfe08dd9eef1f33fca2bd Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Sun, 4 Oct 2026 23:38:31 -0700 Subject: [PATCH 4/4] FIX Close scenario cancellation admission and gather races Stop queue admission on new supervisor cancellation requests and drain the original gather together with all workers. Cover completion callbacks, ready siblings, repeated cancellation, and same-task resume. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/scenario/core/scenario.py | 12 +- .../core/test_scenario_partial_results.py | 196 ++++++++++++++++++ 2 files changed, 205 insertions(+), 3 deletions(-) diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index d629ea720a..eaf666e341 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -1690,6 +1690,8 @@ async def _execute_atomic_attacks_parallel_async( 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. @@ -1717,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: @@ -1762,7 +1770,7 @@ async def worker_async() -> None: for worker in workers: if not worker.done() and not worker.cancelling(): worker.cancel() - drain = asyncio.gather(*workers, return_exceptions=True) + drain = asyncio.gather(group, *workers, return_exceptions=True) caller_cancellation: asyncio.CancelledError | None = None while not drain.done(): try: @@ -1770,8 +1778,6 @@ async def worker_async() -> None: except asyncio.CancelledError as cancellation: caller_cancellation = cancellation drain.result() - # Retrieve an exception that arrived after the initial shield was cancelled. - group.exception() if caller_cancellation is not None: raise caller_cancellation from None raise diff --git a/tests/unit/scenario/core/test_scenario_partial_results.py b/tests/unit/scenario/core/test_scenario_partial_results.py index f3df11f1d5..f04db186c2 100644 --- a/tests/unit/scenario/core/test_scenario_partial_results.py +++ b/tests/unit/scenario/core/test_scenario_partial_results.py @@ -555,6 +555,202 @@ async def run_queued_attack(*args, **kwargs): assert sorted(resumed_result.get_objectives()) == ["obj1", "obj2", "obj3"] assert all(len(results) == 1 for results in resumed_result.attack_results.values()) + @pytest.mark.parametrize("max_retries", [0, 1]) + async def test_run_async_cancellation_waits_for_worker_completion_callbacks_async( + self, *, mock_objective_target: MagicMock, max_retries: int + ) -> None: + attacks = [create_mock_atomic_attack(name, [name]) for name in ("first", "second")] + all_started = asyncio.Event() + release_workers = asyncio.Event() + persisted: set[str] = set() + worker_tasks: list[asyncio.Task[None]] = [] + + def make_run_async(atomic_attack: MagicMock) -> AsyncMock: + async def run_async(**_kwargs: object) -> AttackExecutorResult[AttackResult]: + worker = asyncio.current_task() + assert worker is not None + worker_tasks.append(worker) + name = atomic_attack.atomic_attack_name + result = AttackResult( + conversation_id=f"conv-{name}", + objective=name, + outcome=AttackOutcome.SUCCESS, + executed_turns=1, + ) + await save_attack_results_to_memory_async([result], atomic_attack=atomic_attack) + persisted.add(name) + if len(persisted) == len(attacks): + all_started.set() + await release_workers.wait() + return AttackExecutorResult(completed_results=[result], incomplete_objectives=[]) + + return AsyncMock(side_effect=run_async) + + for attack in attacks: + attack.run_async = make_run_async(attack) + scenario = ConcreteScenario(name="Cancellation Completion Race", version=1, atomic_attacks_to_return=attacks) + scenario.set_params_from_args( + args={"objective_target": mock_objective_target, "max_concurrency": 2, "max_retries": max_retries} + ) + await scenario.initialize_async() + + parent = asyncio.create_task(scenario.run_async()) + try: + await asyncio.wait_for(all_started.wait(), timeout=5) + release_workers.set() + parent.cancel("stop scenario") + with pytest.raises(asyncio.CancelledError, match="stop scenario"): + await asyncio.wait_for(parent, timeout=5) + + assert all(worker.done() for worker in worker_tasks) + assert not scenario._active_atomic_groups + [stored] = await scenario._memory.get_scenario_results_async( + scenario_result_ids=[scenario._scenario_result_id] + ) + assert stored.scenario_run_state is ScenarioRunState.CANCELLED + assert stored.error_type == "CancelledError" + assert stored.number_tries == 1 + assert sorted(stored.get_objectives()) == ["first", "second"] + assert all(attack.run_async.call_count == 1 for attack in attacks) + finally: + release_workers.set() + if not parent.done(): + parent.cancel() + await asyncio.gather(parent, *worker_tasks, return_exceptions=True) + + @pytest.mark.parametrize("cancel_again", [False, True], ids=["single-cancel", "repeated-cancel"]) + async def test_caller_cancellation_stops_queue_before_ready_sibling_finishes_async( + self, *, mock_objective_target: MagicMock, cancel_again: bool + ) -> None: + slow_attack, sibling_attack, queued_attack = [ + create_mock_atomic_attack(name, [name]) for name in ("slow", "sibling", "queued") + ] + all_started = asyncio.Event() + release_sibling = asyncio.Event() + cleanup_started = asyncio.Event() + release_cleanup = asyncio.Event() + cleanup_finished = asyncio.Event() + worker_tasks: list[asyncio.Task[None]] = [] + + async def slow_run_async(**_kwargs: object) -> None: + worker = asyncio.current_task() + assert worker is not None + worker_tasks.append(worker) + try: + await asyncio.Event().wait() + finally: + cleanup_started.set() + await release_cleanup.wait() + cleanup_finished.set() + + async def sibling_run_async(**_kwargs: object) -> AttackExecutorResult[AttackResult]: + worker = asyncio.current_task() + assert worker is not None + worker_tasks.append(worker) + all_started.set() + await release_sibling.wait() + return AttackExecutorResult(completed_results=[], incomplete_objectives=[]) + + slow_attack.run_async = AsyncMock(side_effect=slow_run_async) + sibling_attack.run_async = AsyncMock(side_effect=sibling_run_async) + queued_attack.run_async = AsyncMock( + return_value=AttackExecutorResult(completed_results=[], incomplete_objectives=[]) + ) + scenario = ConcreteScenario( + name="Caller Cancellation Admission", + version=1, + atomic_attacks_to_return=[slow_attack, sibling_attack, queued_attack], + ) + scenario.set_params_from_args( + args={"objective_target": mock_objective_target, "max_concurrency": 2, "max_retries": 2} + ) + await scenario.initialize_async() + + parent = asyncio.create_task(scenario.run_async()) + try: + await asyncio.wait_for(all_started.wait(), timeout=5) + release_sibling.set() + parent.cancel("stop scenario") + await asyncio.wait_for(cleanup_started.wait(), timeout=5) + assert not parent.done() + assert worker_tasks[0].cancelling() == 1 + [stored] = await scenario._memory.get_scenario_results_async( + scenario_result_ids=[scenario._scenario_result_id] + ) + assert stored.scenario_run_state is ScenarioRunState.IN_PROGRESS + + if cancel_again: + parent.cancel("stop scenario again") + await asyncio.sleep(0) + assert not parent.done() + assert worker_tasks[0].cancelling() == 1 + + release_cleanup.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(parent, timeout=5) + assert cleanup_finished.is_set() + assert all(worker.done() for worker in worker_tasks) + assert not scenario._active_atomic_groups + queued_attack.run_async.assert_not_called() + [stored] = await scenario._memory.get_scenario_results_async( + scenario_result_ids=[scenario._scenario_result_id] + ) + assert stored.scenario_run_state is ScenarioRunState.CANCELLED + assert stored.number_tries == 1 + finally: + release_sibling.set() + release_cleanup.set() + if not parent.done(): + parent.cancel() + await asyncio.gather(parent, *worker_tasks, return_exceptions=True) + + async def test_run_async_resumes_in_a_task_with_previous_cancellation_async( + self, *, mock_objective_target: MagicMock + ) -> None: + attack = create_mock_atomic_attack("resumed", ["objective"]) + started = asyncio.Event() + completed_result = AttackResult( + conversation_id="conv-resumed", + objective="objective", + outcome=AttackOutcome.SUCCESS, + executed_turns=1, + ) + + async def run_async(**_kwargs: object) -> AttackExecutorResult[AttackResult]: + if attack.run_async.call_count == 1: + started.set() + await asyncio.Event().wait() + await save_attack_results_to_memory_async([completed_result], atomic_attack=attack) + return AttackExecutorResult(completed_results=[completed_result], incomplete_objectives=[]) + + attack.run_async = AsyncMock(side_effect=run_async) + scenario = ConcreteScenario(name="Resume After Cancellation", version=1, atomic_attacks_to_return=[attack]) + scenario.set_params_from_args(args={"objective_target": mock_objective_target, "max_concurrency": 1}) + await scenario.initialize_async() + + async def cancel_then_resume_async() -> ScenarioResult: + with pytest.raises(asyncio.CancelledError): + await scenario.run_async() + supervisor = asyncio.current_task() + assert supervisor is not None + assert supervisor.cancelling() == 1 + return await scenario.run_async() + + parent = asyncio.create_task(cancel_then_resume_async()) + try: + await asyncio.wait_for(started.wait(), timeout=5) + parent.cancel("stop first run") + result = await asyncio.wait_for(parent, timeout=5) + assert result.scenario_run_state is ScenarioRunState.COMPLETED + assert result.number_tries == 2 + assert result.get_objectives() == ["objective"] + assert attack.run_async.call_count == 2 + assert not scenario._active_atomic_groups + finally: + if not parent.done(): + parent.cancel() + await asyncio.gather(parent, return_exceptions=True) + async def test_worker_cancellation_waits_for_cleanup_despite_caller_cancellation(self, mock_objective_target): cancelled_attack = create_mock_atomic_attack("cancelled", ["obj1"]) sibling_attack = create_mock_atomic_attack("sibling", ["obj2"])