diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 6e8179369b..eaf666e341 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -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. @@ -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: @@ -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) @@ -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) + 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() diff --git a/tests/unit/scenario/core/test_scenario.py b/tests/unit/scenario/core/test_scenario.py index bfef6b0280..f38f83c316 100644 --- a/tests/unit/scenario/core/test_scenario.py +++ b/tests/unit/scenario/core/test_scenario.py @@ -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 ): diff --git a/tests/unit/scenario/core/test_scenario_partial_results.py b/tests/unit/scenario/core/test_scenario_partial_results.py index 05ff591c7a..24b0c36deb 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: @@ -424,8 +426,9 @@ async def mock_run(*args, **kwargs): assert len(result.attack_results["resume_attack"]) == 5 @pytest.mark.timeout(30) + @pytest.mark.parametrize("cancel_worker", [False, True], ids=["caller-cancelled", "worker-cancelled"]) async def test_run_async_cancellation_persists_progress_cleans_workers_and_resumes_async( - self, mock_objective_target: MagicMock + self, *, mock_objective_target: MagicMock, cancel_worker: bool ) -> None: completed_attack = create_mock_atomic_attack("completed_attack", ["obj1"]) in_flight_attack = create_mock_atomic_attack("in_flight_attack", ["obj2"]) @@ -458,18 +461,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): - (await save_attack_results_to_memory_async([completed_result], atomic_attack=completed_attack)) + worker_tasks.append(asyncio.current_task()) + await save_attack_results_to_memory_async([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() @@ -513,13 +520,16 @@ async def run_queued_attack(*args, **kwargs): await scenario_task pytest.fail("Scenario finished before reaching the cancellation checkpoint") await workers_ready - 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() + 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"] @@ -548,7 +558,339 @@ async def run_queued_attack(*args, **kwargs): finally: workers_ready.cancel() scenario_task.cancel() - await asyncio.gather(workers_ready, scenario_task, return_exceptions=True) + await asyncio.gather(workers_ready, scenario_task, *worker_tasks, return_exceptions=True) + + @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"]) + 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) + + @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] = 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") + 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] = 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_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