From df10671c8f6ced77833f20e68026202e9b96883e Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 8 Oct 2026 12:58:53 -0700 Subject: [PATCH 1/2] FIX: Prepare scenarios on the backend event loop Retain serialized preparation and abandoned admission cleanup, offload blocking construction, and preserve runtime shutdown ownership. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 8fabc130-e568-4583-9ff0-ec26491a6414 --- pyrit/backend/services/runtime_lifecycle.py | 3 + .../backend/services/scenario_run_service.py | 289 +++---- .../registry/components/scenario_registry.py | 7 +- pyrit/scenario/scenarios/airt/psychosocial.py | 3 +- .../initializers/load_default_datasets.py | 3 +- .../initializers/preload_scenario_metadata.py | 5 +- pyrit/setup/initializers/scorers.py | 10 + pyrit/setup/initializers/targets.py | 5 + tests/unit/backend/test_runtime_lifecycle.py | 168 +++- tests/unit/backend/test_scenario_resume.py | 60 +- .../unit/backend/test_scenario_run_routes.py | 2 +- .../unit/backend/test_scenario_run_service.py | 747 ++++++++++++------ tests/unit/registry/test_scenario_registry.py | 67 +- tests/unit/scenario/airt/test_psychosocial.py | 46 ++ .../unit/setup/test_load_default_datasets.py | 8 +- .../setup/test_preload_scenario_metadata.py | 13 +- tests/unit/setup/test_scorer_initializer.py | 13 + tests/unit/setup/test_targets_initializer.py | 13 + 18 files changed, 1008 insertions(+), 454 deletions(-) diff --git a/pyrit/backend/services/runtime_lifecycle.py b/pyrit/backend/services/runtime_lifecycle.py index 75632819c4..696a3707e4 100644 --- a/pyrit/backend/services/runtime_lifecycle.py +++ b/pyrit/backend/services/runtime_lifecycle.py @@ -260,6 +260,9 @@ async def shutdown_async(self) -> None: async def _shutdown_runtime_async(self) -> None: """Finish retained requests without cancelling their offloaded writes.""" + service = peek_scenario_run_service() + if service: + service.stop_admission() pending = self.operations | self.management_operations if self.apply_task is not None: pending.add(self.apply_task) diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index fe672e8d59..9f0aa98f0f 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -11,13 +11,11 @@ import asyncio import base64 import contextlib -import functools import json import logging import uuid from collections import OrderedDict, deque from collections.abc import AsyncIterator, Mapping, Sequence -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from datetime import UTC, datetime from typing import Any @@ -39,6 +37,7 @@ ScenarioProgressReadModel, ScenarioProgressSnapshot, ) +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.memory import AttackResultKeysetCursor, CentralMemory, SQLiteMemory from pyrit.memory.memory_interface import ( ScenarioHistoryAggregate, @@ -137,7 +136,7 @@ class ScenarioRunNotFoundError(ValueError): @dataclass class _PreparedRun: - """Scenario and request-scoped default captured in the preparation worker.""" + """Scenario and request-scoped default captured on the backend loop.""" scenario: Scenario adversarial_target: PromptTarget | None = None @@ -178,18 +177,14 @@ class ScenarioRunService: Service for managing scenario run lifecycle. Uses CentralMemory (database) as the source of truth for run state. - Read methods are async-only and run on the backend event loop. + Preparation, execution, and reads run on the backend event loop. + The preparation gate stays held through abandoned-request cleanup. Keeps executable objects in a process-local single-active FIFO scheduler. FIFO ordering therefore spans only runs submitted to the same backend process. Deploy one backend replica to preserve a global admission order; multiple replicas require a shared database-backed scheduler or lease. """ - #: Seconds to let initialization's own background tasks (for example HTTP client teardown - #: scheduled from ``__del__``) finish before the initialization loop is torn down. This is - #: headroom for incidental teardown, not a waiter for real long-running work. - _INITIALIZATION_DRAIN_TIMEOUT = 5.0 - def __init__(self, *, max_concurrent_runs: int = _DEFAULT_MAX_CONCURRENT_RUNS) -> None: """ Initialize the scenario run service. @@ -205,13 +200,8 @@ def __init__(self, *, max_concurrent_runs: int = _DEFAULT_MAX_CONCURRENT_RUNS) - self._progress_read_model = ScenarioProgressReadModel(memory=self._memory) self._technique_metadata_cache: dict[str, dict[str, ScenarioTechniqueSummary]] = {} - # Initialization writes to CentralMemory, and the in-memory SQLite backend shares one - # DBAPI connection across every thread (StaticPool, sqlite_memory.py). Two preparations - # running at once would use that connection concurrently and lose or corrupt writes, so - # they are serialized onto a single worker. The event loop is still free while they run, - # which is the point of the offload. - self._prepare_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="pyrit-scenario-prep") - self._preparations: set[asyncio.Future[_PreparedRun]] = set() + self._preparations: set[asyncio.Task[_PreparedRun]] = set() + self._preparation_errors: list[Exception] = [] self._terminal_errors: OrderedDict[str, str] = OrderedDict() self._active_scenario_result_id: str | None = None self._queued_runs: deque[_ActiveTask] = deque() @@ -219,6 +209,7 @@ def __init__(self, *, max_concurrent_runs: int = _DEFAULT_MAX_CONCURRENT_RUNS) - self._abandoned_prepare_tasks: set[asyncio.Task[None]] = set() self._scheduler_lock = asyncio.Lock() self._launch_lock = asyncio.Lock() + self._preparation_lock = asyncio.Lock() self._preparing_run_ids: set[str] = set() self._pending_resume_requests: set[str] = set() self._queue_revision = 0 @@ -227,14 +218,24 @@ def __init__(self, *, max_concurrent_runs: int = _DEFAULT_MAX_CONCURRENT_RUNS) - def has_active_work(self) -> bool: """Return whether scenario scheduling, preparation, or handoff work remains.""" return bool( - self._active_scenario_result_id or self._queued_runs or self._preparations or self._handoff_retry_tasks + self._active_scenario_result_id + or self._queued_runs + or self._preparations + or self._abandoned_prepare_tasks + or self._handoff_retry_tasks + or self._launch_lock.locked() ) async def close_async(self) -> None: """Close a service only after all tracked work has drained.""" if self.has_active_work(): raise RuntimeError("Scenario work has not drained.") - await asyncio.to_thread(self._prepare_executor.shutdown, wait=True) + if self._preparation_errors: + raise ExceptionGroup("Failed to clean up abandoned scenario preparations.", self._preparation_errors) + + def stop_admission(self) -> None: + """Reject new preparations and prevent queued runs from starting.""" + self._stopping = True async def start_run_async(self, *, request: RunScenarioRequest) -> ScenarioRunSummary: """ @@ -380,29 +381,43 @@ async def _start_run_locked_async(self, *, request: RunScenarioRequest) -> Scena """ if self._stopping: raise RuntimeError("Scenario run scheduling is stopping.") - resumed_from_cancelled = await self._is_run_cancelled_async(scenario_result_id=request.scenario_result_id) - if request.scenario_result_id: - self._preparing_run_ids.add(request.scenario_result_id) - prepare_task = asyncio.get_running_loop().run_in_executor( - self._prepare_executor, - functools.partial(self._prepare_run_blocking, request=request), - ) - self._preparations.add(prepare_task) - prepare_task.add_done_callback(self._discard_preparation) - if request.scenario_result_id: - prepare_task.add_done_callback(lambda _: self._preparing_run_ids.discard(request.scenario_result_id or "")) + await self._preparation_lock.acquire() + prepare_task: asyncio.Task[_PreparedRun] | None = None + abandoned = False try: + if self._stopping: + raise RuntimeError("Scenario run scheduling is stopping.") + resumed_from_cancelled = await self._is_run_cancelled_async(scenario_result_id=request.scenario_result_id) + if request.scenario_result_id: + self._preparing_run_ids.add(request.scenario_result_id) + prepare_task = asyncio.create_task(self._prepare_run_async(request=request)) + self._preparations.add(prepare_task) prepared = await asyncio.shield(prepare_task) + return await self._schedule_prepared_run_async( + request=request, prepared=prepared, resumed_from_cancelled=resumed_from_cancelled + ) except asyncio.CancelledError: - if prepare_task.done(): - try: - (await self._release_abandoned_prepare_async(prepare_task)) - except Exception as cleanup_error: - logger.warning(f"Could not clean up after a cancelled scenario preparation: {cleanup_error}") - else: - prepare_task.add_done_callback(self._schedule_abandoned_prepare_cleanup) + if prepare_task is not None: + abandoned = True + cleanup = asyncio.create_task( + self._cleanup_abandoned_prepare_async(prepare_task=prepare_task, request=request) + ) + self._abandoned_prepare_tasks.add(cleanup) + cleanup.add_done_callback(self._abandoned_prepare_done) raise + finally: + if not abandoned: + self._finish_preparation(prepare_task=prepare_task, request=request) + async def _schedule_prepared_run_async( + self, *, request: RunScenarioRequest, prepared: _PreparedRun, resumed_from_cancelled: bool + ) -> ScenarioRunSummary: + """ + Validate persisted preparation and admit it to the FIFO scheduler. + + Returns: + ScenarioRunSummary: The persisted active, queued, or terminal state. + """ scenario = prepared.scenario scenario_result_id = scenario._scenario_result_id if scenario_result_id is None: @@ -456,8 +471,15 @@ async def _start_run_locked_async(self, *, request: RunScenarioRequest) -> Scena raise RuntimeError(f"Scenario run {scenario_result_id} was not found in the database after initialization.") return response - def _discard_preparation(self, preparation: asyncio.Future[_PreparedRun]) -> None: - self._preparations.discard(preparation) + def _finish_preparation( + self, *, prepare_task: asyncio.Task[_PreparedRun] | None, request: RunScenarioRequest + ) -> None: + """Release preparation ownership only after admission or abandoned cleanup.""" + if prepare_task is not None: + self._preparations.discard(prepare_task) + if request.scenario_result_id: + self._preparing_run_ids.discard(request.scenario_result_id) + self._preparation_lock.release() async def _is_run_cancelled_async(self, *, scenario_result_id: str | None) -> bool: """ @@ -477,140 +499,54 @@ async def _is_run_cancelled_async(self, *, scenario_result_id: str | None) -> bo stored = await self._memory.get_scenario_result_header_async(scenario_result_id=scenario_result_id) return stored is not None and stored.scenario_run_state == ScenarioRunState.CANCELLED - async def _release_abandoned_prepare_async(self, prepare_task: "asyncio.Future[_PreparedRun]") -> None: + async def _cleanup_abandoned_prepare_async( + self, *, prepare_task: asyncio.Task[_PreparedRun], request: RunScenarioRequest + ) -> None: + """Retain the preparation gate and resume reservation through cleanup.""" + try: + await self._release_abandoned_prepare_async(prepare_task) + finally: + self._finish_preparation(prepare_task=prepare_task, request=request) + + def _abandoned_prepare_done(self, task: asyncio.Task[None]) -> None: + """Retrieve cleanup failures and retain them for runtime shutdown.""" + self._abandoned_prepare_tasks.discard(task) + if not task.cancelled() and (error := task.exception()) is not None: + logger.error("Could not clean up a cancelled scenario preparation.", exc_info=error) + if isinstance(error, Exception): + self._preparation_errors.append(error) + + async def _release_abandoned_prepare_async(self, prepare_task: asyncio.Task[_PreparedRun]) -> None: """ - Clean up after an abandoned preparation thread has finished. + Wait for abandoned preparation and conditionally cancel its stored result. - A preparation that succeeds after its caller is cancelled leaves behind a - scenario result nobody will run. Mark it cancelled rather than leaving it - in ``CREATED``. + Keep scheduler-owned runs intact. If cancellation interrupts preparation + or admission, terminalize its unowned nonterminal result. Args: - prepare_task: The future wrapping the abandoned ``_prepare_run_blocking`` call. + prepare_task: The retained backend-loop preparation. """ - if prepare_task.cancelled(): - return - error = prepare_task.exception() - if error is not None: - logger.warning(f"Abandoned scenario preparation failed after the request was cancelled: {error}") + try: + prepared = await prepare_task + except Exception as error: + logger.warning("Abandoned scenario preparation failed after the request was cancelled: %s", error) return - # Initialization already stored a CREATED scenario result, and nothing is going to run - # it now, so terminalize it rather than leaving a run that never starts. A run that - # already reached a terminal state keeps it, so a real failure is not relabelled. - scenario_result_id = prepare_task.result().scenario._scenario_result_id + scenario_result_id = prepared.scenario._scenario_result_id + if scenario_result_id in self._active_tasks or any( + run.scenario_result_id == scenario_result_id for run in self._queued_runs + ): + return if scenario_result_id: - try: - ( - await self._memory.try_update_scenario_run_state_async( - scenario_result_id=scenario_result_id, - expected_states={ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS}, - scenario_run_state=ScenarioRunState.CANCELLED, - error_message="The start request was cancelled while the scenario was being initialized.", - ) - ) - except Exception as update_error: - logger.warning( - f"Could not mark abandoned scenario run {scenario_result_id} as cancelled: {update_error}" - ) + await self._memory.try_update_scenario_run_state_async( + scenario_result_id=scenario_result_id, + # Admission can persist QUEUED before cancellation prevents the queue append. + expected_states={ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS, ScenarioRunState.QUEUED}, + scenario_run_state=ScenarioRunState.CANCELLED, + error_message="The start request was cancelled before the scenario was admitted to the scheduler.", + ) logger.warning("Abandoned scenario preparation completed after the request was cancelled.") - def _schedule_abandoned_prepare_cleanup(self, prepare_task: "asyncio.Future[_PreparedRun]") -> None: - """Keep async cleanup alive after a preparation request is cancelled.""" - task = asyncio.create_task(self._release_abandoned_prepare_async(prepare_task)) - self._abandoned_prepare_tasks.add(task) - task.add_done_callback(self._abandoned_prepare_tasks.discard) - - def _prepare_run_blocking(self, *, request: RunScenarioRequest) -> _PreparedRun: - """ - Run the eager initialization for a scenario run on the calling thread. - - Exists so ``start_run_async`` can offload initialization onto a worker thread. - The scenario is executed later on the caller's event loop, so initialization must not - leave anything bound to the throwaway loop used here. Clients that schedule their own - teardown are given a moment to finish; anything still running after that would be - cancelled when the loop closes, so the start fails rather than handing back a scenario - that holds dead async resources. - - Args: - request: The run request with scenario name, target, and options. - - Returns: - _PreparedRun: The initialized scenario and its scoped default. - - Raises: - RuntimeError: If tasks are still running on the initialization loop after the drain. - """ - - async def prepare_and_drain_async() -> _PreparedRun: - prepared = await self._prepare_run_async(request=request) - try: - await self._drain_initialization_tasks_async() - except RuntimeError as drain_error: - # Initialization already stored a CREATED row and this start is over, so - # terminalize it here rather than leaving a run that never begins. A cancel - # can land while the drain is running, so keep whatever terminal state won. - scenario_result_id = prepared.scenario._scenario_result_id - if scenario_result_id: - try: - ( - await self._memory.try_update_scenario_run_state_async( - scenario_result_id=scenario_result_id, - expected_states={ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS}, - scenario_run_state=ScenarioRunState.FAILED, - error_message=str(drain_error), - error_type=type(drain_error).__name__, - ) - ) - except Exception as update_error: - logger.warning(f"Could not mark scenario run {scenario_result_id} as failed: {update_error}") - raise - return prepared - - async def prepare_async() -> _PreparedRun: - try: - return await prepare_and_drain_async() - finally: - await self._memory.dispose_loop_resources_async() - - return asyncio.run(prepare_async()) - - async def _drain_initialization_tasks_async(self) -> None: - """ - Let initialization's background tasks finish before the initialization loop closes. - - Initialization builds throwaway async clients, and some of them schedule their own - teardown from ``__del__``, so a task can appear purely because a garbage collection - landed late. Waiting for those is the difference between a scenario that starts and - one that fails at random. A draining task can also start another one, so the set is - rebuilt after every wait and the whole drain shares a single deadline. - - Raises: - RuntimeError: If any task is still running after the drain timeout. - """ - loop = asyncio.get_running_loop() - current_task = asyncio.current_task() - deadline = loop.time() + self._INITIALIZATION_DRAIN_TIMEOUT - - while True: - pending = [task for task in asyncio.all_tasks() if task is not current_task] - if not pending: - return - - remaining = deadline - loop.time() - if remaining <= 0: - raise RuntimeError( - "Scenario initialization left background tasks on the initialization loop, which is " - "about to close. They would be cancelled and the scenario would hold dead async " - f"resources: {', '.join(sorted(task.get_name() for task in pending))}" - ) - - done, _ = await asyncio.wait(pending, timeout=remaining) - for task in done: - # Retrieve outcomes so a failed teardown task does not log "never retrieved" noise. - if not task.cancelled() and task.exception() is not None: - logger.debug(f"A scenario initialization task failed during teardown: {task.exception()}") - async def _prepare_run_async(self, *, request: RunScenarioRequest) -> _PreparedRun: """ Resolve and initialize the scenario for a run request. @@ -624,14 +560,17 @@ async def _prepare_run_async(self, *, request: RunScenarioRequest) -> _PreparedR Raises: ValueError: If scenario, target, initializer, or technique cannot be found. """ - scenario_class = self._configuration_resolver.resolve_scenario_class(scenario_name=request.scenario_name) + scenario_class = await run_legacy_sync_async( + self._configuration_resolver.resolve_scenario_class, scenario_name=request.scenario_name + ) await self._run_initializers_async(request=request) objective_target = self._configuration_resolver.resolve_target(target_name=request.target_name) adversarial_target = self._configuration_resolver.resolve_adversarial_target( target_name=request.adversarial_target_name ) with override_default_adversarial_target(adversarial_target): - init_kwargs = self._configuration_resolver.resolve_configuration( + init_kwargs = await run_legacy_sync_async( + self._configuration_resolver.resolve_configuration, scenario_name=request.scenario_name, scenario_class=scenario_class, objective_target=objective_target, @@ -934,12 +873,12 @@ async def reconcile_interrupted_runs_async(self) -> int: async def shutdown_async(self) -> None: """Stop scheduling and terminalize active and queued runs for process shutdown.""" + self.stop_admission() task: asyncio.Task[None] | None = None retry_tasks: list[asyncio.Task[None]] = [] errors: list[Exception] = [] async with self._launch_lock: async with self._scheduler_lock: - self._stopping = True retry_tasks = list(self._handoff_retry_tasks) queued = list(self._queued_runs) self._queued_runs.clear() @@ -974,9 +913,10 @@ async def shutdown_async(self) -> None: self._active_scenario_result_id = None self._release_completed_task(scenario_result_id=active.scenario_result_id) self._queue_revision += 1 - await asyncio.to_thread(self._prepare_executor.shutdown, wait=True) if self._abandoned_prepare_tasks: - await asyncio.gather(*self._abandoned_prepare_tasks) + await asyncio.gather(*self._abandoned_prepare_tasks, return_exceptions=True) + errors.extend(self._preparation_errors) + self._preparation_errors.clear() if task is not None and not task.done(): task.cancel() try: @@ -996,7 +936,14 @@ async def _enqueue_run_async(self, *, scheduled: _ActiveTask) -> None: """Atomically enqueue a persisted initialized run or start it immediately.""" async with self._scheduler_lock: if self._stopping: - raise RuntimeError("Scenario run scheduling is stopping.") + await self._memory.try_update_scenario_run_state_async( + scenario_result_id=scheduled.scenario_result_id, + expected_states={ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS}, + scenario_run_state=ScenarioRunState.FAILED, + error_message=_SHUTDOWN_INTERRUPTION_REASON, + error_type=_INTERRUPTED_ERROR_TYPE, + ) + return scheduled_ids = { *(run.scenario_result_id for run in self._queued_runs), *self._active_tasks.keys(), @@ -1193,12 +1140,12 @@ async def _run_initializers_async(self, *, request: RunScenarioRequest) -> None: if not request.initializers: return - initializer_registry = InitializerRegistry.get_registry_singleton() + initializer_registry = await run_legacy_sync_async(InitializerRegistry.get_registry_singleton) for initializer_name in request.initializers: initializer_params = (request.initializer_args or {}).get(initializer_name) try: - instance = initializer_registry.create_and_configure( - initializer_name, initializer_params=initializer_params + instance = await run_legacy_sync_async( + initializer_registry.create_and_configure, initializer_name, initializer_params=initializer_params ) except KeyError as e: raise ValueError(f"Initializer not found: {e}") from None diff --git a/pyrit/registry/components/scenario_registry.py b/pyrit/registry/components/scenario_registry.py index a11afd7828..ca95a6a724 100644 --- a/pyrit/registry/components/scenario_registry.py +++ b/pyrit/registry/components/scenario_registry.py @@ -18,6 +18,7 @@ from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Literal +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.models import ScenarioRunSizeEstimate, ScenarioTechniqueSummary, class_name_to_snake_case from pyrit.models.identifiers.scenario_identifier import ScenarioIdentifier from pyrit.registry.registry import ParamBagRegistry @@ -271,6 +272,8 @@ async def create_and_initialize_async( Prefer this over manually chaining ``create_instance`` + ``set_params_from_args`` + ``initialize_async``. + Synchronous construction and configuration run off-loop; async + initialization runs on the caller's event loop. Args: name (str): The registry name of the scenario (e.g. ``"foundry.red_team_agent"``). @@ -291,7 +294,9 @@ async def create_and_initialize_async( constructor_kwargs["scenario_result_id"] = scenario_result_id merged_args = {**(scenario_params or {}), **initialize_kwargs} - scenario = self._create_and_configure(name, params=merged_args, constructor_kwargs=constructor_kwargs) + scenario = await run_legacy_sync_async( + self._create_and_configure, name, params=merged_args, constructor_kwargs=constructor_kwargs + ) scenario.set_scenario_registry_name(scenario_registry_name=name) if initial_metadata: scenario.set_initial_metadata(metadata=initial_metadata) diff --git a/pyrit/scenario/scenarios/airt/psychosocial.py b/pyrit/scenario/scenarios/airt/psychosocial.py index 37e68e6781..9d600801ba 100644 --- a/pyrit/scenario/scenarios/airt/psychosocial.py +++ b/pyrit/scenario/scenarios/airt/psychosocial.py @@ -640,7 +640,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list for technique in techniques: if technique is PsychosocialTechnique.Crescendo: - attack_technique = self._build_crescendo_technique( + attack_technique = await asyncio.to_thread( + self._build_crescendo_technique, harm=harm, objective_target=context.objective_target, adversarial_chat=adversarial_chat, diff --git a/pyrit/setup/initializers/load_default_datasets.py b/pyrit/setup/initializers/load_default_datasets.py index bdaf053490..0600eba5b5 100644 --- a/pyrit/setup/initializers/load_default_datasets.py +++ b/pyrit/setup/initializers/load_default_datasets.py @@ -12,6 +12,7 @@ import logging import textwrap +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.datasets import SeedDatasetFilter, SeedDatasetProvider from pyrit.memory import CentralMemory from pyrit.models.parameter import Parameter @@ -80,7 +81,7 @@ async def initialize_async(self) -> None: unique_datasets = list(dict.fromkeys(matched)) logger.info(f"Loading {len(unique_datasets)} dataset(s) matching tags: {sorted(tags)}") else: - unique_datasets = self._scenario_default_dataset_names() + unique_datasets = await run_legacy_sync_async(self._scenario_default_dataset_names) logger.info(f"Loading {len(unique_datasets)} unique datasets required by all scenarios") if not unique_datasets: diff --git a/pyrit/setup/initializers/preload_scenario_metadata.py b/pyrit/setup/initializers/preload_scenario_metadata.py index bddd0072dd..687aaaf4aa 100644 --- a/pyrit/setup/initializers/preload_scenario_metadata.py +++ b/pyrit/setup/initializers/preload_scenario_metadata.py @@ -13,6 +13,7 @@ import logging +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.registry import ScenarioRegistry from pyrit.setup.pyrit_initializer import PyRITInitializer @@ -24,6 +25,6 @@ class PreloadScenarioMetadata(PyRITInitializer): async def initialize_async(self) -> None: """Warm the scenario metadata cache.""" - registry = ScenarioRegistry.get_registry_singleton() - metadata = registry.get_all_registered_class_metadata() + registry = await run_legacy_sync_async(ScenarioRegistry.get_registry_singleton) + metadata = await run_legacy_sync_async(registry.get_all_registered_class_metadata) logger.info("Preloaded metadata for %d scenarios", len(metadata)) diff --git a/pyrit/setup/initializers/scorers.py b/pyrit/setup/initializers/scorers.py index 43b96f1f05..e8e8dd31bf 100644 --- a/pyrit/setup/initializers/scorers.py +++ b/pyrit/setup/initializers/scorers.py @@ -23,6 +23,7 @@ from azure.ai.contentsafety.models import TextCategory +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.models import SeedPrompt from pyrit.models.parameter import Parameter from pyrit.registry import ScorerRegistry, TargetRegistry @@ -187,6 +188,15 @@ async def initialize_async(self) -> None: Raises: RuntimeError: If the TargetRegistry is empty or hasn't been initialized. """ + await run_legacy_sync_async(self._register_scorers) + + def _register_scorers(self) -> None: + """ + Construct template-backed scorers without blocking the caller's event loop. + + Raises: + RuntimeError: If the target registry is empty. + """ target_registry = TargetRegistry.get_registry_singleton() if len(target_registry.instances) == 0: diff --git a/pyrit/setup/initializers/targets.py b/pyrit/setup/initializers/targets.py index 60ec173365..52e031ad6f 100644 --- a/pyrit/setup/initializers/targets.py +++ b/pyrit/setup/initializers/targets.py @@ -20,6 +20,7 @@ from typing import Any from pyrit.auth import get_azure_openai_auth, get_azure_token_provider +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.models.identifiers import TARGET_EVAL_PARAM_FALLBACKS, TARGET_EVAL_PARAMS from pyrit.models.parameter import Parameter from pyrit.prompt_target import ( @@ -625,6 +626,10 @@ async def initialize_async(self) -> None: and target class are automatically grouped into ``RoundRobinTarget`` instances for rate-limit distribution and fault tolerance. """ + await run_legacy_sync_async(self._register_targets) + + def _register_targets(self) -> None: + """Construct and register targets without blocking the caller's event loop.""" tags = self.params.get("tags", ["default"]) if TargetInitializerTags.ALL in tags: tags = [tag for tag in TargetInitializerTags if tag != TargetInitializerTags.ALL] diff --git a/tests/unit/backend/test_runtime_lifecycle.py b/tests/unit/backend/test_runtime_lifecycle.py index 2bd671b654..51cccc2170 100644 --- a/tests/unit/backend/test_runtime_lifecycle.py +++ b/tests/unit/backend/test_runtime_lifecycle.py @@ -22,9 +22,13 @@ from pyrit.backend.services.configuration_file_service import ConfigurationFileService from pyrit.backend.services.manual_send_scheduler import get_manual_send_scheduler from pyrit.backend.services.runtime_lifecycle import RuntimeLifecycle -from pyrit.backend.services.scenario_run_service import ScenarioRunService +from pyrit.backend.services.scenario_run_service import ScenarioRunService, _PreparedRun from pyrit.memory import CentralMemory, MemoryInterface +from pyrit.models import ScenarioRunState +from pyrit.models.catalog.scenario import RunScenarioRequest +from pyrit.scenario import Scenario from pyrit.setup.configuration_loader import ConfigurationLoader +from unit.mocks import make_scenario_result @pytest.fixture @@ -277,6 +281,168 @@ async def test_background_estimates_reject_apply(runtime: RuntimeLifecycle) -> N lifecycle_module.close_services_async.assert_not_awaited() +async def test_http_disconnect_retains_scenario_launch_async(runtime: RuntimeLifecycle) -> None: + """HTTP cancellation leaves the middleware-owned service waiter alive.""" + memory = MagicMock(spec=MemoryInterface) + record = make_scenario_result(attack_results={}, scenario_run_state=ScenarioRunState.CREATED) + memory.get_scenario_results_async.return_value = [record] + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + service = ScenarioRunService() + scenario = MagicMock(spec=Scenario) + scenario._scenario_result_id = str(record.id) + scenario.active_atomic_group_ids = set() + entered, release, executed = asyncio.Event(), asyncio.Event(), asyncio.Event() + + async def prepare_async(*, request: RunScenarioRequest) -> _PreparedRun: + entered.set() + await release.wait() + return _PreparedRun(scenario=scenario) + + async def run_async() -> None: + executed.set() + record.scenario_run_state = ScenarioRunState.COMPLETED + + scenario.run_async = AsyncMock(side_effect=run_async) + + @runtime.app.post("/api/scenarios/runs") + async def launch_async(request: RunScenarioRequest) -> dict[str, str]: + result = await service.start_run_async(request=request) + return {"scenario_result_id": result.scenario_result_id} + + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async), + patch.object(lifecycle_module, "peek_scenario_run_service", return_value=service), + ): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=runtime.app), base_url="http://test") as client: + request = asyncio.create_task( + client.post("/api/scenarios/runs", json={"scenario_name": "test", "target_name": "test"}) + ) + try: + await asyncio.wait_for(entered.wait(), 5) + request.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(request, 5) + assert runtime.operations + assert service.has_active_work() + assert not service._abandoned_prepare_tasks + await apply_async(runtime) + assert runtime.outcome == "busy" + release.set() + await asyncio.wait_for(asyncio.gather(*runtime.operations), 5) + await asyncio.wait_for(executed.wait(), 5) + scenario.run_async.assert_awaited_once() + assert not service._abandoned_prepare_tasks + finally: + release.set() + await asyncio.gather(request, *runtime.operations, return_exceptions=True) + await asyncio.wait_for(service.shutdown_async(), 5) + + +@pytest.mark.parametrize("phase", ["preparation", "cleanup"]) +@pytest.mark.parametrize("cleanup_fails", [False, True]) +@pytest.mark.parametrize("cancel_shutdown", [False, True]) +async def test_runtime_retains_abandoned_preparation_until_memory_shutdown_async( + *, runtime: RuntimeLifecycle, phase: str, cleanup_fails: bool, cancel_shutdown: bool +) -> None: + memory = MagicMock(spec=MemoryInterface) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + service = ScenarioRunService() + scenario = MagicMock(spec=Scenario) + scenario._scenario_result_id = "abandoned" + entered, release, cleanup_entered, cleanup_release, stopped = (asyncio.Event() for _ in range(5)) + order: list[str] = [] + original_stop = service.stop_admission + + def stop() -> None: + original_stop() + stopped.set() + + async def prepare_async(*, request: RunScenarioRequest) -> _PreparedRun: + entered.set() + await release.wait() + order.append("prepared") + return _PreparedRun(scenario=scenario) + + async def cleanup_async(**kwargs: object) -> bool: + assert kwargs["expected_states"] == { + ScenarioRunState.CREATED, + ScenarioRunState.IN_PROGRESS, + ScenarioRunState.QUEUED, + } + cleanup_entered.set() + await cleanup_release.wait() + order.append("cleaned") + if cleanup_fails: + raise RuntimeError("cleanup persistence failed") + return True + + async def dispose_async() -> None: + assert not service.has_active_work() + order.append("memory-closed") + + memory.try_update_scenario_run_state_async.side_effect = cleanup_async + memory.dispose_engine_async.side_effect = dispose_async + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async) as prepare, + patch.object(service, "stop_admission", side_effect=stop), + patch.object(lifecycle_module, "peek_scenario_run_service", return_value=service), + patch.object(CentralMemory, "_memory_instance", memory), + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + ): + start = asyncio.create_task( + service.start_run_async(request=RunScenarioRequest(scenario_name="test", target_name="test")) + ) + shutdown: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(entered.wait(), 5) + await apply_async(runtime) + assert runtime.outcome == "busy" + start.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(start, 5) + if phase == "cleanup": + release.set() + await asyncio.wait_for(cleanup_entered.wait(), 5) + await apply_async(runtime) + assert runtime.outcome == "busy" + lifecycle_module.close_services_async.assert_not_awaited() + shutdown = asyncio.create_task(runtime.shutdown_async()) + await asyncio.wait_for(stopped.wait(), 5) + assert service._stopping + assert not shutdown.done() + memory.dispose_engine_async.assert_not_awaited() + with pytest.raises(RuntimeError, match="scheduling is stopping"): + await service.start_run_async(request=RunScenarioRequest(scenario_name="test", target_name="test")) + prepare.assert_awaited_once() + if cancel_shutdown: + shutdown.cancel() + barrier = asyncio.Event() + asyncio.get_running_loop().call_soon(barrier.set) + await asyncio.wait_for(barrier.wait(), 5) + assert not shutdown.done() + assert all(task.cancelling() == 0 for task in service._preparations) + release.set() + await asyncio.wait_for(cleanup_entered.wait(), 5) + assert not shutdown.done() + memory.dispose_engine_async.assert_not_awaited() + cleanup_release.set() + if cleanup_fails: + with pytest.raises(ExceptionGroup, match="shutdown transitions") as error: + await asyncio.wait_for(shutdown, 5) + assert str(error.value.exceptions[0]) == "cleanup persistence failed" + elif cancel_shutdown: + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(shutdown, 5) + else: + await asyncio.wait_for(shutdown, 5) + assert order == ["prepared", "cleaned", "memory-closed"] + scenario.run_async.assert_not_called() + finally: + release.set() + cleanup_release.set() + await asyncio.gather(start, *([shutdown] if shutdown else []), return_exceptions=True) + + async def test_apply_denies_writes_but_allows_repair_reads(runtime: RuntimeLifecycle) -> None: async with runtime.edit_lock: async with httpx.AsyncClient(transport=httpx.ASGITransport(app=runtime.app), base_url="http://test") as client: diff --git a/tests/unit/backend/test_scenario_resume.py b/tests/unit/backend/test_scenario_resume.py index 472dc09042..fa6014e131 100644 --- a/tests/unit/backend/test_scenario_resume.py +++ b/tests/unit/backend/test_scenario_resume.py @@ -20,7 +20,7 @@ ) from pyrit.exceptions import ScenarioPartialFailureException from pyrit.executor.attack import AttackScoringConfig, PromptSendingAttack -from pyrit.memory import CentralMemory +from pyrit.memory import CentralMemory, SQLiteMemory from pyrit.models import ( SCENARIO_RUN_PLAN_METADATA_KEY, AttackOutcome, @@ -171,6 +171,48 @@ async def fail_second_async(*, normalized_conversation: list[Message]) -> list[M return result +async def test_preparation_and_execution_share_sqlite_loop_resources_async( + resume_environment: tuple[ScenarioRunService, MockPromptTarget], sqlite_instance: SQLiteMemory +) -> None: + service, _ = resume_environment + loop = asyncio.get_running_loop() + original_initialize = _OfflineResumeScenario.initialize_async + original_run = _OfflineResumeScenario.run_async + engines: list[object] = [] + + async def initialize_async(self: _OfflineResumeScenario) -> None: + assert asyncio.get_running_loop() is loop + await original_initialize(self) + engines.append(sqlite_instance._get_async_engine()) + + async def run_async(self: _OfflineResumeScenario) -> None: + assert asyncio.get_running_loop() is loop + engines.append(sqlite_instance._get_async_engine()) + await original_run(self) + + with ( + patch.object(_OfflineResumeScenario, "initialize_async", initialize_async), + patch.object(_OfflineResumeScenario, "run_async", run_async), + patch.object( + sqlite_instance, "dispose_loop_resources_async", wraps=sqlite_instance.dispose_loop_resources_async + ) as dispose, + ): + response = await service.start_run_async( + request=RunScenarioRequest( + scenario_name=_SCENARIO_NAME, + target_name=_TARGET_NAME, + max_concurrency=1, + include_baseline=False, + ) + ) + await _wait_for_idle_async(service) + dispose.assert_not_awaited() + assert len(engines) == 2 and engines[0] is engines[1] + assert sqlite_instance._get_async_engine() is engines[0] + stored = await sqlite_instance.get_scenario_result_header_async(scenario_result_id=response.scenario_result_id) + assert stored is not None and stored.scenario_run_state == ScenarioRunState.COMPLETED + + async def test_resume_preserves_completed_objectives_and_original_id_async( resume_environment: tuple[ScenarioRunService, MockPromptTarget], ) -> None: @@ -242,7 +284,7 @@ async def test_resume_without_launch_metadata_is_rejected_without_initialization stored = await _create_failed_run_async(target=target, legacy=True) run_id = str(stored.id) target.prompt_sent.clear() - with patch.object(service, "_prepare_run_blocking") as prepare: + with patch.object(service, "_prepare_run_async") as prepare: with pytest.raises(ScenarioRunConflictError, match="older run.*cannot be resumed through the GUI"): await service.resume_run_async(scenario_result_id=run_id) prepare.assert_not_called() @@ -304,7 +346,7 @@ async def test_resume_rejects_ineligible_state_before_initializing_async( scenario_result_id=str(stored.id), scenario_run_state=state ) ) - with patch.object(service, "_prepare_run_blocking") as prepare: + with patch.object(service, "_prepare_run_async") as prepare: with pytest.raises(ScenarioRunConflictError, match="cannot resume"): await service.resume_run_async(scenario_result_id=str(stored.id)) prepare.assert_not_called() @@ -428,7 +470,7 @@ async def hold_run_async(*, scenario_result_id: str) -> None: assert resumed.active_scenario_result_id == active.scenario_result_id assert resumed.completed_at is None assert resumed.started_at is None - with patch.object(service, "_prepare_run_blocking") as prepare: + with patch.object(service, "_prepare_run_async") as prepare: with pytest.raises(ScenarioRunConflictError, match="already scheduled"): await service.resume_run_async(scenario_result_id=run_id) prepare.assert_not_called() @@ -460,7 +502,7 @@ async def test_resume_incomplete_saved_configuration_never_uses_defaults_async( scenario_result_id=str(stored.id), metadata=stored.metadata ) ) - with patch.object(service, "_prepare_run_blocking") as prepare: + with patch.object(service, "_prepare_run_async") as prepare: with pytest.raises(ScenarioRunConflictError, match="incomplete"): await service.resume_run_async(scenario_result_id=str(stored.id)) prepare.assert_not_called() @@ -522,7 +564,7 @@ async def test_resume_invalid_saved_configuration_is_rejected_before_initializat stored.metadata[_LAUNCH_REQUEST_METADATA_KEY][field] = value with ( patch.object(service._memory, "get_scenario_result_header_async", return_value=stored), - patch.object(service, "_prepare_run_blocking") as prepare, + patch.object(service, "_prepare_run_async") as prepare, ): with pytest.raises(ScenarioRunConflictError, match="incomplete|invalid|empty"): await service.resume_run_async(scenario_result_id=str(stored.id)) @@ -538,7 +580,7 @@ async def test_resume_missing_canonical_selection_never_uses_current_defaults_as stored.scenario_identifier = stored.scenario_identifier.model_copy(update={missing: None}) with ( patch.object(service._memory, "get_scenario_result_header_async", return_value=stored), - patch.object(service, "_prepare_run_blocking") as prepare, + patch.object(service, "_prepare_run_async") as prepare, ): with pytest.raises(ScenarioRunConflictError, match="missing techniques or datasets"): await service.resume_run_async(scenario_result_id=str(stored.id)) @@ -617,7 +659,7 @@ async def test_original_start_route_also_guards_resume_admission_async( scenario_result_id=str(stored.id), scenario_run_state=state ) ) - with patch.object(service, "_prepare_run_blocking") as prepare: + with patch.object(service, "_prepare_run_async") as prepare: with pytest.raises(ScenarioRunConflictError): await service.start_run_async( request=RunScenarioRequest( @@ -659,7 +701,7 @@ async def test_resume_never_schedules_replacement_result_id_async( stored = await _create_failed_run_async(target=target, legacy=False) replacement = _OfflineResumeScenario(scenario_result_id="different-result-id") with ( - patch.object(service, "_prepare_run_blocking", return_value=_PreparedRun(scenario=replacement)), + patch.object(service, "_prepare_run_async", return_value=_PreparedRun(scenario=replacement)), patch.object(service, "_enqueue_run_async") as enqueue, ): with pytest.raises(ValueError, match="changed the saved result ID"): diff --git a/tests/unit/backend/test_scenario_run_routes.py b/tests/unit/backend/test_scenario_run_routes.py index 63029cbe44..f172251a8b 100644 --- a/tests/unit/backend/test_scenario_run_routes.py +++ b/tests/unit/backend/test_scenario_run_routes.py @@ -228,7 +228,7 @@ def test_older_run_returns_409_without_initialization(self, client: TestClient) service = _svc_mod.ScenarioRunService() with ( patch("pyrit.backend.routes.scenarios.get_scenario_run_service", return_value=service), - patch.object(service, "_prepare_run_blocking") as prepare, + patch.object(service, "_prepare_run_async") as prepare, ): response = client.post(f"/api/scenarios/runs/{stored.id}/resume") prepare.assert_not_called() diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index bc690fa36a..6fb181876f 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -8,7 +8,6 @@ import asyncio import logging import threading -import time import uuid from dataclasses import replace from datetime import UTC, datetime, timedelta @@ -22,6 +21,7 @@ import pyrit.backend.services.scenario_run_service as _svc_mod from pyrit.backend.services.scenario_progress_read_model import ScenarioPlanLookup, ScenarioProgressReadModel from pyrit.backend.services.scenario_run_service import ( + ScenarioRunConflictError, ScenarioRunService, ) from pyrit.common.utils import to_sha256 @@ -117,10 +117,14 @@ async def test_has_active_work_covers_scheduler_owned_work(patch_central_databas assert service.has_active_work() service._queued_runs.clear() - preparation: asyncio.Future[Any] = asyncio.get_running_loop().create_future() + async def prepare_async() -> _svc_mod._PreparedRun: + return _svc_mod._PreparedRun(scenario=MagicMock(spec=Scenario)) + + preparation = asyncio.create_task(prepare_async()) service._preparations.add(preparation) assert service.has_active_work() service._preparations.clear() + await preparation handoff = asyncio.create_task(asyncio.sleep(0)) service._handoff_retry_tasks.add(handoff) @@ -128,6 +132,17 @@ async def test_has_active_work_covers_scheduler_owned_work(patch_central_databas service._handoff_retry_tasks.clear() await handoff + cleanup = asyncio.create_task(asyncio.sleep(0)) + service._abandoned_prepare_tasks.add(cleanup) + assert service.has_active_work() + with pytest.raises(RuntimeError, match="not drained"): + await service.close_async() + service._abandoned_prepare_tasks.clear() + await cleanup + + async with service._launch_lock: + assert service.has_active_work() + assert not service.has_active_work() await service.close_async() @@ -294,7 +309,7 @@ async def test_invalid_explicit_target_never_initializes_async( await service.shutdown_async() @pytest.mark.parametrize("first_outcome", ["success", "error", "cancel"]) - async def test_worker_queue_execution_and_handoff_keep_submitted_targets_async( + async def test_preparation_queue_execution_and_handoff_keep_submitted_targets_async( self, mock_all_registries: dict[str, Any], first_outcome: str ) -> None: targets = {name: MockPromptTarget() for name in ("first", "second", "adversarial_chat", "my_target")} @@ -308,6 +323,7 @@ async def test_worker_queue_execution_and_handoff_keep_submitted_targets_async( release_first = asyncio.Event() finished = asyncio.Event() main_thread = threading.get_ident() + backend_loop = asyncio.get_running_loop() def introspect() -> Any: assert threading.get_ident() != main_thread @@ -315,7 +331,8 @@ def introspect() -> Any: return mock_all_registries["scenario_instance"] async def initialize_async(*args: object, **kwargs: object) -> Any: - assert threading.get_ident() != main_thread + assert threading.get_ident() == main_thread + assert asyncio.get_running_loop() is backend_loop initialized.append(get_default_adversarial_target()) await asyncio.sleep(0) assert get_default_adversarial_target() is initialized[-1] @@ -327,6 +344,7 @@ async def initialize_async(*args: object, **kwargs: object) -> Any: scenario.active_atomic_group_ids = set() async def run_async() -> None: + assert asyncio.get_running_loop() is backend_loop selected = get_default_adversarial_target() executed.append(selected) if run_id == "scope-0": @@ -384,7 +402,7 @@ def update_state(*, scenario_result_id: str, scenario_run_state: ScenarioRunStat release_first.set() await service.shutdown_async() - async def test_failed_preparation_restores_worker_scope_async(self, mock_all_registries: dict[str, Any]) -> None: + async def test_failed_preparation_restores_scope_async(self, mock_all_registries: dict[str, Any]) -> None: selected, fallback = MockPromptTarget(), MockPromptTarget() targets = {"selected": selected, "adversarial_chat": fallback, "my_target": fallback} mock_all_registries["target_registry"].instances.get.side_effect = targets.get @@ -400,30 +418,26 @@ async def fail_async(*args: object, **kwargs: object) -> None: try: with pytest.raises(ValueError, match="initialization failed"): await service.start_run_async(request=request) - worker_default = await asyncio.get_running_loop().run_in_executor( - service._prepare_executor, get_default_adversarial_target - ) - assert worker_default is fallback assert get_default_adversarial_target() is fallback assert not service._active_tasks finally: await service.shutdown_async() - async def test_cancelled_start_keeps_worker_scope_until_preparation_finishes_async( + async def test_cancelled_start_keeps_scope_until_preparation_finishes_async( self, mock_all_registries: dict[str, Any] ) -> None: selected, fallback = MockPromptTarget(), MockPromptTarget() targets = {"selected": selected, "adversarial_chat": fallback, "my_target": fallback} mock_all_registries["target_registry"].instances.get.side_effect = targets.get service = ScenarioRunService() - started = threading.Event() - release = threading.Event() + started = asyncio.Event() + release = asyncio.Event() observed: list[PromptTarget] = [] async def initialize_async(*args: object, **kwargs: object) -> Any: observed.append(get_default_adversarial_target()) started.set() - assert await asyncio.to_thread(release.wait, 5) + await asyncio.wait_for(release.wait(), 5) observed.append(get_default_adversarial_target()) return mock_all_registries["scenario_instance"] @@ -432,16 +446,13 @@ async def initialize_async(*args: object, **kwargs: object) -> Any: request.adversarial_target_name = "selected" task = asyncio.create_task(service.start_run_async(request=request)) try: - assert await asyncio.to_thread(started.wait, 5) + await asyncio.wait_for(started.wait(), 5) task.cancel() with pytest.raises(asyncio.CancelledError): await task assert get_default_adversarial_target() is fallback release.set() - worker_default = await asyncio.get_running_loop().run_in_executor( - service._prepare_executor, get_default_adversarial_target - ) - assert worker_default is fallback + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) assert observed == [selected, selected] assert not service._active_tasks finally: @@ -1097,30 +1108,27 @@ async def test_start_run_omits_scenario_result_id_when_none(self, mock_all_regis assert call.kwargs["scenario_result_id"] is None async def test_start_run_keeps_event_loop_responsive(self, mock_all_registries) -> None: - """Initialization is offloaded, so the loop keeps running while a run starts.""" + """A database await in initialization does not block the backend loop.""" service = ScenarioRunService() + started = asyncio.Event() + release = asyncio.Event() + backend_loop = asyncio.get_running_loop() - def _slow_prepare(*, request: Any) -> Any: - time.sleep(0.5) + async def _slow_prepare_async(*, request: Any) -> Any: + assert asyncio.get_running_loop() is backend_loop + started.set() + await release.wait() return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - beats = 0 - - async def _heartbeat() -> None: - nonlocal beats - while True: - await asyncio.sleep(0.01) - beats += 1 - - with patch.object(service, "_prepare_run_blocking", _slow_prepare): - heartbeat = asyncio.create_task(_heartbeat()) + with patch.object(service, "_prepare_run_async", _slow_prepare_async): + task = asyncio.create_task(service.start_run_async(request=_make_request())) try: - await service.start_run_async(request=_make_request()) + await asyncio.wait_for(started.wait(), 5) + await asyncio.wait_for(asyncio.sleep(0), 5) + assert not task.done() finally: - heartbeat.cancel() - - # A blocked event loop yields zero heartbeats over the same window. - assert beats > 10 + release.set() + await asyncio.wait_for(task, 5) async def test_start_run_background_task_survives_handoff(self, mock_all_registries) -> None: """The background task must outlive start_run_async and actually execute the run.""" @@ -1144,53 +1152,59 @@ async def test_start_run_marks_abandoned_prepare_cancelled(self, mock_all_regist service = ScenarioRunService() scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "abandoned-id" - finished = threading.Event() + started = asyncio.Event() + release = asyncio.Event() - def _slow_prepare(*, request: Any) -> Any: - time.sleep(0.5) - finished.set() + async def _slow_prepare_async(*, request: Any) -> Any: + started.set() + await release.wait() return _svc_mod._PreparedRun(scenario=scenario_instance) - with patch.object(service, "_prepare_run_blocking", _slow_prepare): + with patch.object(service, "_prepare_run_async", _slow_prepare_async): with patch.object(service._memory, "try_update_scenario_run_state_async") as update_state: task = asyncio.create_task(service.start_run_async(request=_make_request())) - await asyncio.sleep(0.1) + await asyncio.wait_for(started.wait(), 5) task.cancel() with pytest.raises(asyncio.CancelledError): await task - await asyncio.sleep(1.0) - assert finished.is_set() + assert service.has_active_work() + release.set() + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + assert not service.has_active_work() + scenario_instance.run_async.assert_not_awaited() update_state.assert_called_once() assert update_state.call_args.kwargs["scenario_result_id"] == "abandoned-id" assert update_state.call_args.kwargs["scenario_run_state"] == ScenarioRunState.CANCELLED assert update_state.call_args.kwargs["expected_states"] == { ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS, + ScenarioRunState.QUEUED, } async def test_start_run_marks_prepare_cancelled_when_it_finishes_before_cancellation_lands( self, mock_all_registries ) -> None: - """A done future never calls back, so this race used to leave the run stuck in CREATED.""" + """Cancellation after preparation finishes must still terminalize the abandoned run.""" service = ScenarioRunService() scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "raced-id" - def _instant_prepare(*, request: Any) -> Any: + async def _instant_prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) - async def _complete_then_cancel(awaitable): + async def _complete_then_cancel_async(awaitable): # The preparation finishes, then the cancellation lands: the exact ordering that # leaves ``prepare_task.done()`` True inside the handler. await awaitable raise asyncio.CancelledError - with patch.object(service, "_prepare_run_blocking", _instant_prepare): - with patch("asyncio.shield", _complete_then_cancel): + with patch.object(service, "_prepare_run_async", _instant_prepare_async): + with patch("asyncio.shield", _complete_then_cancel_async): with patch.object(service._memory, "try_update_scenario_run_state_async") as update_state: with pytest.raises(asyncio.CancelledError): await service.start_run_async(request=_make_request()) + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) update_state.assert_called_once() assert update_state.call_args.kwargs["scenario_result_id"] == "raced-id" @@ -1198,20 +1212,138 @@ async def _complete_then_cancel(awaitable): assert update_state.call_args.kwargs["expected_states"] == { ScenarioRunState.CREATED, ScenarioRunState.IN_PROGRESS, + ScenarioRunState.QUEUED, } + @pytest.mark.parametrize("phase", ["persisted-read", "summary-read", "active-write", "queued-write"]) + async def test_cancellation_before_admission_retains_cleanup_async( + self, mock_all_registries: dict[str, Any], phase: str + ) -> None: + service = ScenarioRunService() + memory = mock_all_registries["memory"] + scenario = mock_all_registries["scenario_instance"] + run_id = scenario._scenario_result_id + record = mock_all_registries["db_result"] + record.scenario_run_state = ScenarioRunState.FAILED + memory.get_scenario_result_header_async.return_value = record + entered, release, cleanup_entered, cleanup_release = (asyncio.Event() for _ in range(4)) + original_response = service._build_response_async + + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + record.scenario_run_state = ScenarioRunState.CREATED + return _svc_mod._PreparedRun(scenario=scenario) + + async def get_results_async(**_: object) -> list[ScenarioResult]: + if phase == "persisted-read": + entered.set() + await release.wait() + return [record] + + async def build_response_async(**kwargs: Any) -> Any: + if phase == "summary-read": + entered.set() + await release.wait() + return await original_response(**kwargs) + + async def update_state_async(*, scenario_run_state: ScenarioRunState, **_: object) -> None: + record.scenario_run_state = scenario_run_state + if phase in {"active-write", "queued-write"}: + entered.set() + await release.wait() + + async def cleanup_async(*, expected_states: set[ScenarioRunState], **_: object) -> bool: + cleanup_entered.set() + await cleanup_release.wait() + assert record.scenario_run_state in expected_states + record.scenario_run_state = ScenarioRunState.CANCELLED + return True + + if phase == "queued-write": + service._active_scenario_result_id = "other-run" + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async), + patch.object(memory, "get_scenario_results_async", side_effect=get_results_async), + patch.object(service, "_build_response_async", side_effect=build_response_async), + patch.object(memory, "update_scenario_run_state_and_metadata_fields_async", side_effect=update_state_async), + patch.object(memory, "try_update_scenario_run_state_async", side_effect=cleanup_async), + ): + start = asyncio.create_task(service.start_run_async(request=_make_request(scenario_result_id=run_id))) + try: + await asyncio.wait_for(entered.wait(), 5) + start.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(start, 5) + await asyncio.wait_for(cleanup_entered.wait(), 5) + assert service.has_active_work() + assert service._preparation_lock.locked() + assert run_id in service._preparing_run_ids + assert not service._active_tasks + assert not service._queued_runs + scenario.run_async.assert_not_awaited() + cleanup_release.set() + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + assert record.scenario_run_state == ScenarioRunState.CANCELLED + assert run_id not in service._preparing_run_ids + assert not service._preparation_lock.locked() + finally: + release.set() + cleanup_release.set() + await asyncio.gather(start, *service._abandoned_prepare_tasks, return_exceptions=True) + service._active_scenario_result_id = None + await service.shutdown_async() + + @pytest.mark.parametrize("queued", [False, True]) + async def test_cancellation_after_admission_preserves_scheduler_ownership_async( + self, mock_all_registries: dict[str, Any], queued: bool + ) -> None: + service = ScenarioRunService() + scenario = mock_all_registries["scenario_instance"] + run_id = scenario._scenario_result_id + entered, release, execution_release = (asyncio.Event() for _ in range(3)) + original_response = service.get_run_from_storage_async + + async def read_response_async(**kwargs: Any) -> Any: + entered.set() + await release.wait() + return await original_response(**kwargs) + + if queued: + service._active_scenario_result_id = "other-run" + scenario.run_async.side_effect = execution_release.wait + with patch.object(service, "get_run_from_storage_async", side_effect=read_response_async): + start = asyncio.create_task(service.start_run_async(request=_make_request())) + try: + await asyncio.wait_for(entered.wait(), 5) + start.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(start, 5) + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + service._memory.try_update_scenario_run_state_async.assert_not_awaited() + if queued: + assert [run.scenario_result_id for run in service._queued_runs] == [run_id] + else: + assert run_id in service._active_tasks + assert not service._preparation_lock.locked() + finally: + release.set() + execution_release.set() + await asyncio.gather(start, return_exceptions=True) + if queued: + service._active_scenario_result_id = None + await service.shutdown_async() + async def test_start_run_does_not_run_a_scenario_cancelled_during_initialization(self, mock_all_registries) -> None: """A run appears in the run list as soon as it is stored, so it can be cancelled mid-init.""" service = ScenarioRunService() scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "cancelled-during-init" - def _prepare(*, request: Any) -> Any: + async def _prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) cancelled = _make_db_scenario_result(result_id="cancelled-during-init", run_state=ScenarioRunState.CANCELLED) mock_all_registries["memory"].get_scenario_results_async = AsyncMock(return_value=[cancelled]) - with patch.object(service, "_prepare_run_blocking", _prepare): + with patch.object(service, "_prepare_run_async", _prepare_async): with patch.object(service, "_execute_run_async") as execute: response = await service.start_run_async(request=_make_request()) @@ -1225,14 +1357,14 @@ async def test_start_run_resumes_a_run_that_was_already_cancelled(self, mock_all scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "resumed-cancelled" - def _prepare(*, request: Any) -> Any: + async def _prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) cancelled = _make_db_scenario_result(result_id="resumed-cancelled", run_state=ScenarioRunState.CANCELLED) mock_all_registries["memory"].get_scenario_results_async = AsyncMock(return_value=[cancelled]) mock_all_registries["memory"].get_scenario_result_header_async = AsyncMock(return_value=cancelled) - with patch.object(service, "_prepare_run_blocking", _prepare): + with patch.object(service, "_prepare_run_async", _prepare_async): with patch.object(service, "_execute_run_async") as execute: response = await service.start_run_async(request=_make_request(scenario_result_id="resumed-cancelled")) @@ -1250,16 +1382,16 @@ async def test_start_run_honours_a_cancel_that_lands_while_a_failed_run_initiali scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "cancelled-mid-init" - def _prepare(*, request: Any) -> Any: + async def _prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) cancelled = _make_db_scenario_result(result_id="cancelled-mid-init", run_state=ScenarioRunState.CANCELLED) failed = _make_db_scenario_result(result_id="cancelled-mid-init", run_state=ScenarioRunState.FAILED) mock_all_registries["memory"].get_scenario_results_async = AsyncMock(return_value=[cancelled]) - # The pre-preparation read sees a live run; the cancel lands while the worker prepares. + # The pre-preparation read sees a live run; cancellation lands during initialization. mock_all_registries["memory"].get_scenario_result_header_async = AsyncMock(return_value=failed) - with patch.object(service, "_prepare_run_blocking", _prepare): + with patch.object(service, "_prepare_run_async", _prepare_async): with patch.object(service, "_execute_run_async") as execute: response = await service.start_run_async(request=_make_request(scenario_result_id="cancelled-mid-init")) @@ -1272,10 +1404,10 @@ async def test_start_run_does_not_read_a_header_for_a_fresh_run(self, mock_all_r scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "fresh-id" - def _prepare(*, request: Any) -> Any: + async def _prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) - with patch.object(service, "_prepare_run_blocking", _prepare): + with patch.object(service, "_prepare_run_async", _prepare_async): with patch.object(service, "_execute_run_async"): await service.start_run_async(request=_make_request()) @@ -1284,10 +1416,10 @@ def _prepare(*, request: Any) -> Any: async def test_start_run_failure_does_not_report_a_cancellation(self, mock_all_registries, caplog) -> None: service = ScenarioRunService() - def _failing_prepare(*, request: Any) -> Any: + async def _failing_prepare_async(*, request: Any) -> Any: raise ValueError("Scenario 'nope' not found") - with patch.object(service, "_prepare_run_blocking", _failing_prepare): + with patch.object(service, "_prepare_run_async", _failing_prepare_async): with caplog.at_level(logging.WARNING): with pytest.raises(ValueError, match="not found"): await service.start_run_async(request=_make_request()) @@ -1295,261 +1427,359 @@ def _failing_prepare(*, request: Any) -> Any: assert "cancelled" not in caplog.text.lower() async def test_start_run_cleanup_failure_still_propagates_cancellation(self, mock_all_registries) -> None: - """Cleanup runs inline on this path, so it must not replace the CancelledError.""" + """Cleanup failures are logged and reported at shutdown, not to the cancelled caller.""" service = ScenarioRunService() scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._scenario_result_id = "raced-id" - def _instant_prepare(*, request: Any) -> Any: + async def _instant_prepare_async(*, request: Any) -> Any: return _svc_mod._PreparedRun(scenario=scenario_instance) - async def _complete_then_cancel(awaitable): + async def _complete_then_cancel_async(awaitable): await awaitable raise asyncio.CancelledError - with patch.object(service, "_prepare_run_blocking", _instant_prepare): - with patch("asyncio.shield", _complete_then_cancel): + with patch.object(service, "_prepare_run_async", _instant_prepare_async): + with patch("asyncio.shield", _complete_then_cancel_async): with patch.object( service, "_release_abandoned_prepare_async", side_effect=RuntimeError("cleanup exploded") ): with pytest.raises(asyncio.CancelledError): await service.start_run_async(request=_make_request()) + await asyncio.gather(*service._abandoned_prepare_tasks, return_exceptions=True) + with pytest.raises(ExceptionGroup, match="shutdown transitions") as caught: + await service.shutdown_async() + assert str(caught.value.exceptions[0]) == "cleanup exploded" - def test_prepare_executor_serializes_preparations(self, mock_all_registries) -> None: - """In-memory SQLite shares one connection across threads, so preparations must not overlap.""" + async def test_start_run_serializes_after_abandoned_prepare(self, mock_all_registries) -> None: + """A cancelled start holds the gate through preparation and result cleanup.""" service = ScenarioRunService() - overlap = [] - active = 0 - lock = threading.Lock() - - def _prepare(*, request: Any) -> Any: - nonlocal active - with lock: - active += 1 - overlap.append(active) - time.sleep(0.05) - with lock: - active -= 1 + started = asyncio.Event() + release = asyncio.Event() + cleanup_started = asyncio.Event() + release_cleanup = asyncio.Event() + second_waiting = asyncio.Event() + second_prepared = asyncio.Event() + preparations = 0 + + async def _prepare_async(*, request: Any) -> Any: + nonlocal preparations + preparations += 1 + if preparations == 1: + started.set() + await release.wait() + else: + second_prepared.set() return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - with patch.object(service, "_prepare_run_blocking", _prepare): - futures = [ - service._prepare_executor.submit(lambda: service._prepare_run_blocking(request=_make_request())) - for _ in range(4) - ] - for future in futures: - future.result() - - assert max(overlap) == 1 - - def test_prepare_run_blocking_waits_for_initialization_teardown_tasks(self, mock_all_registries) -> None: - """Async clients schedule their own teardown, so a benign task must not fail the start.""" - service = ScenarioRunService() + async def _cleanup_async(**kwargs: Any) -> bool: + cleanup_started.set() + await release_cleanup.wait() + return True - async def _prepare_with_teardown(*, request: Any) -> Any: - task = asyncio.create_task(asyncio.sleep(0.05)) - task.set_name("client-teardown-task") - return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) + async def _second_start_async() -> Any: + second_waiting.set() + return await service.start_run_async(request=_make_request()) - with patch.object(service, "_prepare_run_async", _prepare_with_teardown): - prepared = service._prepare_run_blocking(request=_make_request()) - assert prepared.scenario is mock_all_registries["scenario_instance"] + with ( + patch.object(service, "_prepare_run_async", _prepare_async), + patch.object(service._memory, "try_update_scenario_run_state_async", side_effect=_cleanup_async), + ): + first = asyncio.create_task(service.start_run_async(request=_make_request())) + try: + await asyncio.wait_for(started.wait(), 5) + first.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(first, 5) + second = asyncio.create_task(_second_start_async()) + await asyncio.wait_for(second_waiting.wait(), 5) + assert preparations == 1 + release.set() + await asyncio.wait_for(cleanup_started.wait(), 5) + assert service.has_active_work() + assert service._preparation_lock.locked() + assert not second_prepared.is_set() + release_cleanup.set() + await asyncio.wait_for(second, 5) + assert preparations == 2 + finally: + release.set() + release_cleanup.set() + await service.shutdown_async() - def test_prepare_run_blocking_fails_when_a_task_outlives_the_drain(self, mock_all_registries) -> None: - """A task still running when the loop closes is cancelled, so the scenario is unusable.""" + async def test_normal_preparations_are_serialized_async(self, mock_all_registries) -> None: service = ScenarioRunService() - - async def _leaky_prepare(*, request: Any) -> Any: - task = asyncio.create_task(asyncio.sleep(3600)) - task.set_name("stray-initializer-task") - await asyncio.sleep(0) + entered, release, second_waiting = asyncio.Event(), asyncio.Event(), asyncio.Event() + order: list[str] = [] + + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + order.append(request.scenario_name) + if len(order) == 1: + entered.set() + await release.wait() return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - with patch.object(service, "_prepare_run_async", _leaky_prepare): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.05): - with pytest.raises(RuntimeError, match="left background tasks on the initialization loop") as exc_info: - service._prepare_run_blocking(request=_make_request()) - - assert "stray-initializer-task" in str(exc_info.value) - - def test_prepare_run_blocking_marks_the_run_failed_when_the_drain_fails(self, mock_all_registries) -> None: - """Initialization already stored the run, so a failed drain must not leave it in CREATED.""" - service = ScenarioRunService() - scenario_instance = mock_all_registries["scenario_instance"] - scenario_instance._scenario_result_id = "drained-id" - - async def _leaky_prepare(*, request: Any) -> Any: - task = asyncio.create_task(asyncio.sleep(3600)) - task.set_name("stray-initializer-task") - await asyncio.sleep(0) - return _svc_mod._PreparedRun(scenario=scenario_instance) - - with patch.object(service, "_prepare_run_async", _leaky_prepare): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.05): - with patch.object(service._memory, "try_update_scenario_run_state_async") as update_state: - with pytest.raises(RuntimeError, match="left background tasks"): - service._prepare_run_blocking(request=_make_request()) + async def second_async() -> Any: + second_waiting.set() + return await service.start_run_async(request=_make_request(scenario_name="second")) - update_state.assert_called_once() - assert update_state.call_args.kwargs["scenario_result_id"] == "drained-id" - assert update_state.call_args.kwargs["scenario_run_state"] == ScenarioRunState.FAILED - assert update_state.call_args.kwargs["error_type"] == "RuntimeError" - assert update_state.call_args.kwargs["expected_states"] == { - ScenarioRunState.CREATED, - ScenarioRunState.IN_PROGRESS, - } + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async), + patch.object(service, "_schedule_prepared_run_async", new_callable=AsyncMock), + ): + first = asyncio.create_task(service.start_run_async(request=_make_request(scenario_name="first"))) + second: asyncio.Task[Any] | None = None + try: + await asyncio.wait_for(entered.wait(), 5) + second = asyncio.create_task(second_async()) + await asyncio.wait_for(second_waiting.wait(), 5) + assert order == ["first"] + release.set() + await asyncio.wait_for(asyncio.gather(first, second), 5) + assert order == ["first", "second"] + assert not service.has_active_work() + finally: + release.set() + await asyncio.gather(first, *([second] if second else []), return_exceptions=True) - def test_prepare_run_blocking_reports_the_drain_error_when_marking_failed_fails(self, mock_all_registries) -> None: - """A bookkeeping failure must not replace the error that explains the failed start.""" + @pytest.mark.parametrize("stage", ["introspection", "initializer"]) + @pytest.mark.parametrize("cancelled", [False, True]) + async def test_blocking_construction_keeps_backend_responsive_async( + self, mock_all_registries, stage: str, cancelled: bool + ) -> None: service = ScenarioRunService() - scenario_instance = mock_all_registries["scenario_instance"] - scenario_instance._scenario_result_id = "drained-id" - - async def _leaky_prepare(*, request: Any) -> Any: - task = asyncio.create_task(asyncio.sleep(3600)) - task.set_name("stray-initializer-task") - await asyncio.sleep(0) - return _svc_mod._PreparedRun(scenario=scenario_instance) - - with patch.object(service, "_prepare_run_async", _leaky_prepare): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.05): - with patch.object( - service._memory, "try_update_scenario_run_state_async", side_effect=ValueError("gone") - ): - with pytest.raises(RuntimeError, match="left background tasks"): - service._prepare_run_blocking(request=_make_request()) + entered = asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + backend_thread = threading.get_ident() + registry = mock_all_registries["initializer_registry"] + initializer = registry.create_and_configure.return_value + + def construct(*args: Any, **kwargs: Any) -> Any: + assert threading.get_ident() != backend_thread + loop.call_soon_threadsafe(entered.set) + if not release.wait(5): + raise TimeoutError("Construction was not released.") + return initializer if stage == "initializer" else mock_all_registries["scenario_instance"] + + async def initialize_async() -> None: + assert asyncio.get_running_loop() is loop + + initializer.initialize_async.side_effect = initialize_async + constructor = registry.create_and_configure if stage == "initializer" else mock_all_registries["scenario_class"] + constructor.side_effect = construct + task = asyncio.create_task( + service.start_run_async( + request=_make_request( + initializers=["target"] if stage == "initializer" else None, + max_dataset_size=1 if stage == "introspection" else None, + ) + ) + ) + try: + await asyncio.wait_for(entered.wait(), 5) + assert service.has_active_work() + assert not task.done() + if cancelled: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + assert service.has_active_work() + assert service._preparation_lock.locked() + release.set() + if cancelled: + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + mock_all_registries["scenario_instance"].run_async.assert_not_awaited() + else: + await asyncio.wait_for(task, 5) + if stage == "initializer": + initializer.initialize_async.assert_awaited_once() + finally: + release.set() + await asyncio.gather(task, return_exceptions=True) + await service.shutdown_async() - def test_prepare_run_blocking_waits_for_a_task_spawned_during_the_drain(self, mock_all_registries) -> None: - """A draining task can start another one, which the first snapshot never saw.""" + @pytest.mark.parametrize( + "state", + [ScenarioRunState.CREATED, ScenarioRunState.FAILED, ScenarioRunState.CANCELLED, ScenarioRunState.COMPLETED], + ) + async def test_abandoned_success_preserves_terminal_state_async( + self, mock_all_registries, state: ScenarioRunState + ) -> None: service = ScenarioRunService() - scenario_instance = mock_all_registries["scenario_instance"] - child_finished = threading.Event() - - async def _prepare_spawning_a_child(*, request: Any) -> Any: - async def _child() -> None: - await asyncio.sleep(0.05) - child_finished.set() - - async def _parent() -> None: - await asyncio.sleep(0.01) - asyncio.create_task(_child(), name="child-teardown-task") + entered, release = asyncio.Event(), asyncio.Event() + record = _make_db_scenario_result(run_state=state) - asyncio.create_task(_parent(), name="parent-teardown-task") - await asyncio.sleep(0) - return _svc_mod._PreparedRun(scenario=scenario_instance) + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + entered.set() + await release.wait() + return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - with patch.object(service, "_prepare_run_async", _prepare_spawning_a_child): - assert service._prepare_run_blocking(request=_make_request()).scenario is scenario_instance + async def update_async( + *, expected_states: set[ScenarioRunState], scenario_run_state: ScenarioRunState, **_: Any + ) -> bool: + if record.scenario_run_state not in expected_states: + return False + record.scenario_run_state = scenario_run_state + return True - assert child_finished.is_set() + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async), + patch.object(service._memory, "try_update_scenario_run_state_async", side_effect=update_async), + ): + task = asyncio.create_task(service.start_run_async(request=_make_request())) + await asyncio.wait_for(entered.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + release.set() + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + assert record.scenario_run_state == (ScenarioRunState.CANCELLED if state == ScenarioRunState.CREATED else state) + mock_all_registries["scenario_instance"].run_async.assert_not_awaited() + assert not service.has_active_work() - def test_prepare_run_blocking_fails_when_a_spawned_child_outlives_the_drain(self, mock_all_registries) -> None: - """The child is the one that would be cancelled by the closing loop, so it must be named.""" + async def test_abandoned_failure_is_retrieved_and_logged_async(self, mock_all_registries, caplog) -> None: service = ScenarioRunService() + entered, release = asyncio.Event(), asyncio.Event() - async def _prepare_spawning_a_slow_child(*, request: Any) -> Any: - async def _parent() -> None: - await asyncio.sleep(0.01) - asyncio.create_task(asyncio.sleep(3600), name="child-teardown-task") - - asyncio.create_task(_parent(), name="parent-teardown-task") - await asyncio.sleep(0) - return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + entered.set() + await release.wait() + raise ValueError("failed retained initialization") - with patch.object(service, "_prepare_run_async", _prepare_spawning_a_slow_child): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.3): - with pytest.raises(RuntimeError, match="child-teardown-task"): - service._prepare_run_blocking(request=_make_request()) + with patch.object(service, "_prepare_run_async", side_effect=prepare_async): + task = asyncio.create_task(service.start_run_async(request=_make_request())) + await asyncio.wait_for(entered.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + release.set() + await asyncio.wait_for(asyncio.gather(*service._abandoned_prepare_tasks), 5) + assert "failed retained initialization" in caplog.text + service._memory.try_update_scenario_run_state_async.assert_not_awaited() + assert not service.has_active_work() - def test_prepare_run_blocking_drain_uses_one_deadline_across_generations(self, mock_all_registries) -> None: - """Re-scanning must not restart the budget, or a chain of tasks could stall a start forever.""" + async def test_duplicate_resume_rejected_until_abandoned_cleanup_finishes_async(self, mock_all_registries) -> None: service = ScenarioRunService() + entered, release, cleanup_entered, release_cleanup = (asyncio.Event() for _ in range(4)) + run_id = "sr-uuid-1" + service._memory.get_scenario_result_header_async.return_value = _make_db_scenario_result( + result_id=run_id, run_state=ScenarioRunState.FAILED + ) - async def _prepare_spawning_a_chain(*, request: Any) -> Any: - async def _link(depth: int) -> None: - await asyncio.sleep(0.05) - asyncio.create_task(_link(depth + 1), name=f"chain-task-{depth + 1}") - - asyncio.create_task(_link(0), name="chain-task-0") - await asyncio.sleep(0) + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + entered.set() + await release.wait() return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - started = time.monotonic() - with patch.object(service, "_prepare_run_async", _prepare_spawning_a_chain): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.3): - with pytest.raises(RuntimeError, match="chain-task"): - service._prepare_run_blocking(request=_make_request()) + async def cleanup_async(**kwargs: Any) -> bool: + cleanup_entered.set() + await release_cleanup.wait() + return True - assert time.monotonic() - started < 3 + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async) as prepare, + patch.object(service._memory, "try_update_scenario_run_state_async", side_effect=cleanup_async), + ): + task = asyncio.create_task(service.start_run_async(request=_make_request(scenario_result_id=run_id))) + try: + await asyncio.wait_for(entered.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + release.set() + await asyncio.wait_for(cleanup_entered.wait(), 5) + assert run_id in service._preparing_run_ids + assert service.has_active_work() + with pytest.raises(ScenarioRunConflictError, match="scheduled or initializing"): + await service.start_run_async(request=_make_request(scenario_result_id=run_id)) + with pytest.raises(ScenarioRunConflictError, match="scheduled or initializing"): + await service.resume_run_async(scenario_result_id=run_id) + prepare.assert_awaited_once() + finally: + release.set() + release_cleanup.set() + await asyncio.wait_for(service.shutdown_async(), 5) + assert run_id not in service._preparing_run_ids + assert not service.has_active_work() - async def test_start_run_fails_when_initialization_leaks_a_task(self, mock_all_registries) -> None: - """A preparation with leaked tasks must fail before scheduling.""" + async def test_waiting_preparation_rechecks_shutdown_admission_async(self, mock_all_registries) -> None: service = ScenarioRunService() + entered, release, second_entered, stopped = (asyncio.Event() for _ in range(4)) + original_stop = service.stop_admission - async def _leaky_prepare(*, request: Any) -> Any: - task = asyncio.create_task(asyncio.sleep(3600)) - task.set_name("stray-initializer-task") - await asyncio.sleep(0) + async def prepare_async(*, request: RunScenarioRequest) -> _svc_mod._PreparedRun: + entered.set() + await release.wait() return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - with patch.object(service, "_prepare_run_async", _leaky_prepare): - with patch.object(ScenarioRunService, "_INITIALIZATION_DRAIN_TIMEOUT", 0.05): - with pytest.raises(RuntimeError, match="left background tasks"): - await service.start_run_async(request=_make_request()) + async def second_async() -> Any: + second_entered.set() + return await service.start_run_async(request=_make_request()) - def test_prepare_run_blocking_is_quiet_when_initialization_is_self_contained( - self, mock_all_registries, caplog - ) -> None: - """The happy path must not warn, otherwise the signal is worthless.""" - service = ScenarioRunService() + def stop() -> None: + original_stop() + stopped.set() - async def _clean_prepare(*, request: Any) -> Any: - return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - - with patch.object(service, "_prepare_run_async", _clean_prepare): - with caplog.at_level(logging.WARNING): - service._prepare_run_blocking(request=_make_request()) - - assert "left background tasks" not in caplog.text + with ( + patch.object(service, "_prepare_run_async", side_effect=prepare_async) as prepare, + patch.object(service, "stop_admission", side_effect=stop), + ): + first = asyncio.create_task(service.start_run_async(request=_make_request())) + await asyncio.wait_for(entered.wait(), 5) + first.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(first, 5) + second = asyncio.create_task(second_async()) + await asyncio.wait_for(second_entered.wait(), 5) + shutdown = asyncio.create_task(service.shutdown_async()) + try: + await asyncio.wait_for(stopped.wait(), 5) + assert not second.done() + assert not shutdown.done() + release.set() + with pytest.raises(RuntimeError, match="scheduling is stopping"): + await asyncio.wait_for(second, 5) + await asyncio.wait_for(shutdown, 5) + prepare.assert_awaited_once() + assert not service.has_active_work() + finally: + release.set() + await asyncio.gather(first, second, shutdown, return_exceptions=True) - async def test_start_run_serializes_after_abandoned_prepare(self, mock_all_registries) -> None: - """A cancelled start must leave its worker isolated until initialization finishes.""" + async def test_initialization_tasks_stay_on_live_backend_loop(self, mock_all_registries) -> None: + """Preparation does not drain or cancel unrelated tasks on the live loop.""" service = ScenarioRunService() - started = threading.Event() - release = threading.Event() - finished = threading.Event() + release = asyncio.Event() + child_started = asyncio.Event() + children: list[asyncio.Task[None]] = [] - def _blocking_prepare(*, request: Any) -> Any: - started.set() - release.wait() - finished.set() + async def _child_async() -> None: + child_started.set() + await release.wait() + + async def _prepare_async(*, request: Any) -> Any: + children.append(asyncio.create_task(_child_async())) return _svc_mod._PreparedRun(scenario=mock_all_registries["scenario_instance"]) - with patch.object(service, "_prepare_run_blocking", _blocking_prepare): - task = asyncio.create_task(service.start_run_async(request=_make_request())) + with patch.object(service, "_prepare_run_async", _prepare_async): try: - assert await asyncio.to_thread(started.wait, 5) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(task, timeout=5) - - assert not finished.is_set() + await asyncio.wait_for(service.start_run_async(request=_make_request()), 5) + await asyncio.wait_for(child_started.wait(), 5) + assert len(children) == 1 and not children[0].done() finally: - task.cancel() release.set() - await asyncio.gather(task, return_exceptions=True) + await asyncio.wait_for(asyncio.gather(*children), 5) await service.shutdown_async() - assert finished.is_set() - async def test_start_run_propagates_prepare_failure(self, mock_all_registries) -> None: """Preparation failures must reach the caller.""" service = ScenarioRunService() - def _failing_prepare(*, request: Any) -> Any: + async def _failing_prepare_async(*, request: Any) -> Any: raise ValueError("boom") - with patch.object(service, "_prepare_run_blocking", _failing_prepare): + with patch.object(service, "_prepare_run_async", _failing_prepare_async): with pytest.raises(ValueError, match="boom"): await service.start_run_async(request=_make_request()) @@ -2484,8 +2714,8 @@ async def test_shutdown_waits_for_preparation_before_terminalizing_run(self, moc ), } active_started = asyncio.Event() - preparation_started = threading.Event() - release_preparation = threading.Event() + preparation_started = asyncio.Event() + release_preparation = asyncio.Event() shutdown_started = asyncio.Event() async def _run_active() -> None: @@ -2519,13 +2749,12 @@ def _try_update_state( record.scenario_run_state = scenario_run_state return True - def _prepare(*, request: Any) -> _svc_mod._PreparedRun: + async def _prepare_async(*, request: Any) -> _svc_mod._PreparedRun: preparation_started.set() - if not release_preparation.wait(timeout=5): - raise TimeoutError("Test did not release scenario preparation.") + await asyncio.wait_for(release_preparation.wait(), 5) return _svc_mod._PreparedRun(scenario=prepared_scenario) - async def _shutdown() -> None: + async def _shutdown_async() -> None: shutdown_started.set() await service.shutdown_async() @@ -2544,10 +2773,10 @@ async def _shutdown() -> None: active.task = asyncio.create_task(service._execute_run_async(scenario_result_id="active")) await active_started.wait() - with patch.object(service, "_prepare_run_blocking", _prepare): + with patch.object(service, "_prepare_run_async", _prepare_async): start_task = asyncio.create_task(service.start_run_async(request=_make_request())) - assert await asyncio.to_thread(preparation_started.wait, 5) - shutdown_task = asyncio.create_task(_shutdown()) + await asyncio.wait_for(preparation_started.wait(), 5) + shutdown_task = asyncio.create_task(_shutdown_async()) await shutdown_started.wait() assert not shutdown_task.done() diff --git a/tests/unit/registry/test_scenario_registry.py b/tests/unit/registry/test_scenario_registry.py index 6722782574..3defa6dcd6 100644 --- a/tests/unit/registry/test_scenario_registry.py +++ b/tests/unit/registry/test_scenario_registry.py @@ -3,12 +3,21 @@ """Tests for ScenarioRegistry._build_metadata and create_and_initialize_async.""" -from unittest.mock import AsyncMock, MagicMock +import asyncio +import threading +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pyrit.registry.components.scenario_registry import ScenarioRegistry -from pyrit.scenario.core import BaselineAttackPolicy, ScenarioTechnique +from pyrit.scenario import Scenario +from pyrit.scenario.core import ( + BaselineAttackPolicy, + ScenarioTechnique, + get_default_adversarial_target, + override_default_adversarial_target, +) +from unit.mocks import MockPromptTarget class _NotNoArgScenario: @@ -75,6 +84,60 @@ class _MarkdownMetadataScenario(_MetadataScenario): """ +@pytest.mark.parametrize("cancelled", [False, True]) +async def test_construction_offload_retains_context_and_ownership_async( + *, patch_central_database: object, cancelled: bool +) -> None: + registry = ScenarioRegistry() + scenario = MagicMock(spec=Scenario) + scenario.initialize_async = AsyncMock() + selected = MockPromptTarget() + loop = asyncio.get_running_loop() + backend_thread = threading.get_ident() + entered = asyncio.Event() + release = threading.Event() + configured: list[object] = [] + + def construct(name: str, *, params: dict, constructor_kwargs: dict) -> MagicMock: + assert threading.get_ident() != backend_thread + assert get_default_adversarial_target() is selected + configured.append(params["objective_target"]) + loop.call_soon_threadsafe(entered.set) + if not release.wait(5): + raise TimeoutError("Construction was not released.") + return scenario + + async def initialize_async() -> None: + assert asyncio.get_running_loop() is loop + assert get_default_adversarial_target() is selected + + scenario.initialize_async.side_effect = initialize_async + with patch.object(registry, "_create_and_configure", side_effect=construct): + with override_default_adversarial_target(selected): + task = asyncio.create_task(registry.create_and_initialize_async("test", objective_target=selected)) + try: + await asyncio.wait_for(entered.wait(), 5) + if cancelled: + task.cancel() + barrier = asyncio.Event() + loop.call_soon(barrier.set) + await asyncio.wait_for(barrier.wait(), 5) + assert not task.done() + release.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + scenario.initialize_async.assert_not_awaited() + else: + assert not task.done() + release.set() + assert await asyncio.wait_for(task, 5) is scenario + scenario.initialize_async.assert_awaited_once() + assert configured == [selected] + finally: + release.set() + await asyncio.gather(task, return_exceptions=True) + + def test_build_metadata_raises_when_scenario_requires_constructor_args() -> None: """Scenarios that cannot be instantiated with no args must surface a clear error.""" registry = ScenarioRegistry() diff --git a/tests/unit/scenario/airt/test_psychosocial.py b/tests/unit/scenario/airt/test_psychosocial.py index 1bb1f2c438..d0e7ee0cdd 100644 --- a/tests/unit/scenario/airt/test_psychosocial.py +++ b/tests/unit/scenario/airt/test_psychosocial.py @@ -3,6 +3,9 @@ """Tests for the Psychosocial scenario (per-sub-harm simulated crescendo swept across converters).""" +import asyncio +import threading +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -24,6 +27,7 @@ ) from pyrit.prompt_target import PromptTarget from pyrit.registry import TargetRegistry +from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.dataset_configuration import ( CompoundDatasetAttackConfiguration, DatasetAttackConfiguration, @@ -121,6 +125,48 @@ def register_default_targets(): FIXTURES = ["patch_central_database"] +@pytest.mark.usefixtures(*FIXTURES) +async def test_crescendo_construction_keeps_initialization_responsive_async( + mock_objective_target: PromptTarget, +) -> None: + scenario = _scenario_with_mock_scorers() + scenario.set_params_from_args( + args={ + "objective_target": mock_objective_target, + "scenario_techniques": [PsychosocialTechnique.Crescendo], + "sub_harm": "imminent_crisis", + "include_baseline": False, + } + ) + loop = asyncio.get_running_loop() + backend_thread = threading.get_ident() + entered = asyncio.Event() + release = threading.Event() + original_build = scenario._build_crescendo_technique + + def build(**kwargs: Any) -> AttackTechnique: + loop.call_soon_threadsafe(entered.set) + assert threading.get_ident() != backend_thread + if not release.wait(5): + raise TimeoutError("Crescendo construction was not released.") + return original_build(**kwargs) + + with ( + _patch_base_seed_groups(_make_seed_groups()), + patch.object(scenario, "_build_crescendo_technique", side_effect=build), + ): + initialize = asyncio.create_task(scenario.initialize_async()) + try: + await asyncio.wait_for(entered.wait(), 5) + assert not initialize.done() + release.set() + await asyncio.wait_for(initialize, 5) + finally: + release.set() + await asyncio.gather(initialize, return_exceptions=True) + assert [attack.atomic_attack_name for attack in scenario._atomic_attacks] == ["imminent_crisis_crescendo"] + + @pytest.mark.usefixtures(*FIXTURES) @pytest.mark.parametrize("outer_limit", [3, None]) @pytest.mark.parametrize("sub_harm", ["all", "imminent_crisis"]) diff --git a/tests/unit/setup/test_load_default_datasets.py b/tests/unit/setup/test_load_default_datasets.py index 120a301df8..91ddfc898f 100644 --- a/tests/unit/setup/test_load_default_datasets.py +++ b/tests/unit/setup/test_load_default_datasets.py @@ -6,6 +6,7 @@ """ from dataclasses import dataclass, field +from threading import get_ident from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -63,8 +64,13 @@ def test_required_env_vars_property(self) -> None: async def test_initialize_async_no_scenarios(self) -> None: """Test initialization when no scenarios are registered.""" initializer = LoadDefaultDatasets() + backend_thread = get_ident() - with patch.object(ScenarioRegistry, "get_all_registered_class_metadata", return_value=[]): + def get_metadata() -> list[_FakeMetadata]: + assert get_ident() != backend_thread + return [] + + with patch.object(ScenarioRegistry, "get_all_registered_class_metadata", side_effect=get_metadata): with patch.object(SeedDatasetProvider, "fetch_datasets_async", new_callable=AsyncMock) as mock_fetch: with patch.object(CentralMemory, "get_memory_instance") as mock_memory: mock_memory_instance = MagicMock() diff --git a/tests/unit/setup/test_preload_scenario_metadata.py b/tests/unit/setup/test_preload_scenario_metadata.py index 7278087517..039f4748d3 100644 --- a/tests/unit/setup/test_preload_scenario_metadata.py +++ b/tests/unit/setup/test_preload_scenario_metadata.py @@ -3,6 +3,7 @@ """Tests for the PreloadScenarioMetadata initializer.""" +from threading import get_ident from unittest.mock import MagicMock, patch import pytest @@ -18,11 +19,13 @@ async def test_initialize_async_warms_metadata_cache(self) -> None: initializer = PreloadScenarioMetadata() mock_registry = MagicMock() - mock_registry.get_all_registered_class_metadata.return_value = [ - MagicMock(), - MagicMock(), - MagicMock(), - ] + backend_thread = get_ident() + + def get_metadata() -> list[MagicMock]: + assert get_ident() != backend_thread + return [MagicMock(), MagicMock(), MagicMock()] + + mock_registry.get_all_registered_class_metadata.side_effect = get_metadata with patch( "pyrit.setup.initializers.preload_scenario_metadata.ScenarioRegistry.get_registry_singleton", diff --git a/tests/unit/setup/test_scorer_initializer.py b/tests/unit/setup/test_scorer_initializer.py index 3137e59c2f..96c4ca26d1 100644 --- a/tests/unit/setup/test_scorer_initializer.py +++ b/tests/unit/setup/test_scorer_initializer.py @@ -2,6 +2,7 @@ # Licensed under the MIT license. import os +from threading import get_ident from unittest.mock import MagicMock, patch import pytest @@ -23,6 +24,18 @@ ) +async def test_scorer_registration_runs_off_loop_async() -> None: + initializer = ScorerInitializer() + backend_thread = get_ident() + + def register() -> None: + assert get_ident() != backend_thread + + with patch.object(initializer, "_register_scorers", side_effect=register) as registration: + await initializer.initialize_async() + registration.assert_called_once() + + class TestScorerInitializerBasic: """Tests for ScorerInitializer class - basic functionality.""" diff --git a/tests/unit/setup/test_targets_initializer.py b/tests/unit/setup/test_targets_initializer.py index 295c99b207..a09220b549 100644 --- a/tests/unit/setup/test_targets_initializer.py +++ b/tests/unit/setup/test_targets_initializer.py @@ -2,6 +2,7 @@ # Licensed under the MIT license. import os +from threading import get_ident from unittest.mock import patch import pytest @@ -12,6 +13,18 @@ from pyrit.setup.initializers.targets import TARGET_CONFIGS, _auto_group_enabled, generate_rr_name, get_behavioral_key +async def test_target_registration_runs_off_loop_async() -> None: + initializer = TargetInitializer() + backend_thread = get_ident() + + def register() -> None: + assert get_ident() != backend_thread + + with patch.object(initializer, "_register_targets", side_effect=register) as registration: + await initializer.initialize_async() + registration.assert_called_once() + + @pytest.mark.parametrize( ("value", "expected"), [(False, False), (True, True), ("false", False), ("YES", True), (["no"], False), (["true"], True)], From bc93337ea61d5c0a58ca194f47093d3a00678468 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 9 Oct 2026 18:17:32 -0700 Subject: [PATCH 2/2] FIX: Keep synchronous scenario attack construction off-loop Offload matrix, template, and converter construction while keeping async dataset reads and persistence on the backend loop. Cover a real seven-dataset RapidResponse launch and template reads for ManyShot, Crescendo, TAP, and Skeleton Key. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 8fabc130-e568-4583-9ff0-ec26491a6414 --- .../core/matrix_atomic_attack_builder.py | 3 + pyrit/scenario/scenarios/airt/cyber.py | 4 +- pyrit/scenario/scenarios/airt/jailbreak.py | 23 +++- pyrit/scenario/scenarios/airt/leakage.py | 4 +- pyrit/scenario/scenarios/airt/multilingual.py | 10 ++ pyrit/scenario/scenarios/airt/psychosocial.py | 8 +- .../scenario/scenarios/airt/rapid_response.py | 4 +- .../scenarios/benchmark/adversarial.py | 4 +- pyrit/scenario/scenarios/garak/doctor.py | 4 +- tests/unit/backend/test_scenario_resume.py | 123 +++++++++++++++++- .../unit/scenario/airt/test_rapid_response.py | 48 +++++++ 11 files changed, 219 insertions(+), 16 deletions(-) diff --git a/pyrit/scenario/core/matrix_atomic_attack_builder.py b/pyrit/scenario/core/matrix_atomic_attack_builder.py index bcb3a4cd4d..f9b2823708 100644 --- a/pyrit/scenario/core/matrix_atomic_attack_builder.py +++ b/pyrit/scenario/core/matrix_atomic_attack_builder.py @@ -188,6 +188,9 @@ def build_matrix_atomic_attacks( Scenarios needing extra axes (adversarial targets, caching, converter stacks) call ``MatrixAtomicAttackBuilder`` directly instead. + This synchronous builder can read attack templates from disk. Async callers + must offload construction, not their dataset reads or persistence. + Args: context (ScenarioContext): The resolved runtime inputs for this run. Supplies the objective target, memory labels, per-dataset seed groups, selected techniques, and diff --git a/pyrit/scenario/scenarios/airt/cyber.py b/pyrit/scenario/scenarios/airt/cyber.py index 2f622c7b41..bb7d2bab54 100644 --- a/pyrit/scenario/scenarios/airt/cyber.py +++ b/pyrit/scenario/scenarios/airt/cyber.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.common.path import SCORER_SEED_PROMPT_PATH from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration from pyrit.scenario.core.matrix_atomic_attack_builder import build_matrix_atomic_attacks @@ -123,7 +124,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list Returns: list[AtomicAttack]: The generated atomic attacks. """ - return build_matrix_atomic_attacks( + return await run_legacy_sync_async( + build_matrix_atomic_attacks, context=context, objective_scorer=self._objective_scorer, technique_converters=self._technique_converters, diff --git a/pyrit/scenario/scenarios/airt/jailbreak.py b/pyrit/scenario/scenarios/airt/jailbreak.py index 9f8e1472d8..ce514595d9 100644 --- a/pyrit/scenario/scenarios/airt/jailbreak.py +++ b/pyrit/scenario/scenarios/airt/jailbreak.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, ClassVar from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.converter import TextJailbreakConverter from pyrit.datasets import TextJailBreak from pyrit.executor.attack.single_turn.prompt_sending import PromptSendingAttack @@ -304,12 +305,14 @@ async def _resolve_templates_async(self) -> list[str]: " or `jailbreak_names` (specific selection)." ) if jailbreak_names: - available = set(TextJailBreak.get_jailbreak_templates()) + available = set(await run_legacy_sync_async(TextJailBreak.get_jailbreak_templates)) diff = set(jailbreak_names) - available if diff: raise ValueError(f"Error: could not find templates `{diff}`!") return list(jailbreak_names) - return TextJailBreak.get_jailbreak_templates(num_templates=num_jailbreaks or _DEFAULT_NUM_JAILBREAKS) + return await run_legacy_sync_async( + TextJailBreak.get_jailbreak_templates, num_templates=num_jailbreaks or _DEFAULT_NUM_JAILBREAKS + ) def _build_initial_scenario_metadata(self) -> dict[str, Any]: """ @@ -461,6 +464,18 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list ) self._resolved_jailbreaks = await self._resolve_templates_async() + return await run_legacy_sync_async(self._build_atomic_attacks, context=context) + + def _build_atomic_attacks(self, *, context: ScenarioContext) -> list[AtomicAttack]: + """ + Build the synchronous template and delivery matrix off-loop. + + Returns: + list[AtomicAttack]: The attacks for the selected templates and delivery methods. + + Raises: + ValueError: If only system-prompt delivery is selected for an incompatible target. + """ num_attempts = self.params["num_jailbreak_attempts"] technique_factories = resolve_technique_factories(context=context, extra_factories=_extra_default_factories()) @@ -468,7 +483,7 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list prompt_sending_factory = technique_factories.get(_PROMPT_SENDING) system_selected = _JAILBREAK_SYSTEM_PROMPT in technique_factories - build_system_delivery = system_selected and self._target_supports_system_delivery(self._objective_target) + build_system_delivery = system_selected and self._target_supports_system_delivery(context.objective_target) if system_selected and not build_system_delivery: if prompt_sending_factory is None: raise ValueError( @@ -481,7 +496,7 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list ) builder = MatrixAtomicAttackBuilder( - objective_target=self._objective_target, + objective_target=context.objective_target, objective_scorer=self._objective_scorer, memory_labels=context.memory_labels, ) diff --git a/pyrit/scenario/scenarios/airt/leakage.py b/pyrit/scenario/scenarios/airt/leakage.py index 6df8c914a8..16446e207b 100644 --- a/pyrit/scenario/scenarios/airt/leakage.py +++ b/pyrit/scenario/scenarios/airt/leakage.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.common.path import SCORER_SEED_PROMPT_PATH from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration @@ -139,7 +140,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list Returns: list[AtomicAttack]: The generated atomic attacks. """ - return build_matrix_atomic_attacks( + return await run_legacy_sync_async( + build_matrix_atomic_attacks, context=context, objective_scorer=self._objective_scorer, technique_converters=self._technique_converters, diff --git a/pyrit/scenario/scenarios/airt/multilingual.py b/pyrit/scenario/scenarios/airt/multilingual.py index b48ffbe124..bd76337f34 100644 --- a/pyrit/scenario/scenarios/airt/multilingual.py +++ b/pyrit/scenario/scenarios/airt/multilingual.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Literal from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.common.path import DATASETS_PATH from pyrit.converter import RandomTranslationConverter, TranslationConverter from pyrit.executor.attack import PromptSendingAttack @@ -306,6 +307,15 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list ) self._resolved_languages = await self._resolve_languages_async() + return await run_legacy_sync_async(self._build_atomic_attacks, context=context) + + def _build_atomic_attacks(self, *, context: ScenarioContext) -> list[AtomicAttack]: + """ + Build the synchronous attack and converter matrix off-loop. + + Returns: + list[AtomicAttack]: The attacks for the selected languages and translation methods. + """ adversarial_chat = self._adversarial_chat or get_default_adversarial_target() strategies = set(self.params.get("translation_strategies") or [_TRANSLATION, _RANDOM_TRANSLATION]) technique_factories = resolve_technique_factories( diff --git a/pyrit/scenario/scenarios/airt/psychosocial.py b/pyrit/scenario/scenarios/airt/psychosocial.py index e539a4eee5..a2d494aa58 100644 --- a/pyrit/scenario/scenarios/airt/psychosocial.py +++ b/pyrit/scenario/scenarios/airt/psychosocial.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, ClassVar, cast from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.common.path import DATASETS_PATH from pyrit.converter import ( CharSwapConverter, @@ -650,11 +651,14 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list max_turns=max_turns, ) else: - converter = _converter_for_technique(technique, adversarial_chat=adversarial_chat) + converter = await run_legacy_sync_async( + _converter_for_technique, technique, adversarial_chat=adversarial_chat + ) extra_converters = ( ConverterConfiguration.from_converters(converters=[converter]) if converter else None ) - attack_technique = base_factory.create( + attack_technique = await run_legacy_sync_async( + base_factory.create, objective_target=context.objective_target, attack_scoring_config=scoring_config, adversarial_chat=adversarial_chat, diff --git a/pyrit/scenario/scenarios/airt/rapid_response.py b/pyrit/scenario/scenarios/airt/rapid_response.py index 4fd292bbe8..6b737cd42e 100644 --- a/pyrit/scenario/scenarios/airt/rapid_response.py +++ b/pyrit/scenario/scenarios/airt/rapid_response.py @@ -17,6 +17,7 @@ from typing import TYPE_CHECKING from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration from pyrit.scenario.core.matrix_atomic_attack_builder import build_matrix_atomic_attacks from pyrit.scenario.core.scenario import Scenario @@ -124,7 +125,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list Returns: list[AtomicAttack]: The generated atomic attacks. """ - return build_matrix_atomic_attacks( + return await run_legacy_sync_async( + build_matrix_atomic_attacks, context=context, objective_scorer=self._objective_scorer, display_group_fn=lambda combo: combo.dataset_name, diff --git a/pyrit/scenario/scenarios/benchmark/adversarial.py b/pyrit/scenario/scenarios/benchmark/adversarial.py index ba7033ee19..699724f896 100644 --- a/pyrit/scenario/scenarios/benchmark/adversarial.py +++ b/pyrit/scenario/scenarios/benchmark/adversarial.py @@ -14,6 +14,7 @@ from pyrit.analytics import get_cached_results_for_technique_async from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.common.path import EXECUTOR_SEED_PROMPT_PATH, SCORER_SEED_PROMPT_PATH from pyrit.common.utils import to_sha256 from pyrit.models import ( @@ -522,7 +523,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list # ``--adversarial-targets`` so per-model ASR rolls up naturally — not any internal # field on the PromptTarget instance (e.g. ``_model_name``). The builder's default # ``{technique}__{target}_{dataset}`` naming preserves the VERSION=2 cache key shape. - atomic_attacks = builder.build( + atomic_attacks = await run_legacy_sync_async( + builder.build, technique_factories=technique_factories, dataset_groups=context.seed_groups_by_dataset, adversarial_targets=resolved_targets, diff --git a/pyrit/scenario/scenarios/garak/doctor.py b/pyrit/scenario/scenarios/garak/doctor.py index 82239b63e1..df9003b17a 100644 --- a/pyrit/scenario/scenarios/garak/doctor.py +++ b/pyrit/scenario/scenarios/garak/doctor.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, ClassVar from pyrit.common import apply_defaults +from pyrit.common.async_compatibility import run_legacy_sync_async from pyrit.converter import LeetspeakConverter, PolicyPuppetryConverter, PolicyPuppetryTemplate from pyrit.executor.attack import AttackConverterConfig, PromptSendingAttack from pyrit.prompt_normalizer import ConverterConfiguration @@ -168,7 +169,8 @@ async def _build_atomic_attacks_async(self, *, context: ScenarioContext) -> list objective_scorer=self._objective_scorer, memory_labels=context.memory_labels, ) - return builder.build( + return await run_legacy_sync_async( + builder.build, technique_factories=technique_factories, dataset_groups=context.seed_groups_by_dataset, include_baseline=context.include_baseline, diff --git a/tests/unit/backend/test_scenario_resume.py b/tests/unit/backend/test_scenario_resume.py index fa6014e131..2d803a17eb 100644 --- a/tests/unit/backend/test_scenario_resume.py +++ b/tests/unit/backend/test_scenario_resume.py @@ -4,8 +4,9 @@ """Offline resume coverage using real scenario persistence and harmless mocked targets.""" import asyncio -from collections.abc import AsyncIterator -from typing import ClassVar +import threading +from collections.abc import AsyncIterator, Sequence +from typing import Any, ClassVar from unittest.mock import AsyncMock, patch import pytest @@ -19,7 +20,8 @@ _PreparedRun, ) from pyrit.exceptions import ScenarioPartialFailureException -from pyrit.executor.attack import AttackScoringConfig, PromptSendingAttack +from pyrit.executor.attack import AttackScoringConfig, ManyShotJailbreakAttack, PromptSendingAttack +from pyrit.executor.attack.single_turn import many_shot_jailbreak from pyrit.memory import CentralMemory, SQLiteMemory from pyrit.models import ( SCENARIO_RUN_PLAN_METADATA_KEY, @@ -28,13 +30,23 @@ Parameter, ScenarioResult, ScenarioRunState, + Seed, SeedObjective, ) from pyrit.models.catalog.scenario import RunScenarioRequest -from pyrit.registry import ScenarioRegistry, TargetRegistry +from pyrit.registry import AttackTechniqueRegistry, ScenarioRegistry, TargetRegistry from pyrit.scenario import DatasetAttackConfiguration -from pyrit.scenario.core import AtomicAttack, AttackTechnique, BaselineAttackPolicy, Scenario, ScenarioTechnique +from pyrit.scenario.core import ( + AtomicAttack, + AttackTechnique, + BaselineAttackPolicy, + Scenario, + ScenarioTechnique, + get_default_adversarial_target, +) +from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory from pyrit.scenario.core.scenario_context import ScenarioContext +from pyrit.scenario.scenarios.airt.rapid_response import RapidResponse, _build_rapid_response_technique from pyrit.score import SubStringScorer from unit.mocks import MockPromptTarget @@ -128,6 +140,107 @@ async def wait_async() -> None: await asyncio.wait_for(wait_async(), timeout=10) +async def test_real_matrix_launch_keeps_heartbeat_alive_during_disk_reads_async( + patch_central_database: object, sqlite_instance: SQLiteMemory +) -> None: + loop = asyncio.get_running_loop() + backend_thread = threading.get_ident() + entered, heartbeat, stop = (asyncio.Event() for _ in range(3)) + release = threading.Event() + original_load = many_shot_jailbreak.load_many_shot_jailbreaking_dataset + reads = 0 + target = MockPromptTarget() + scorer = SubStringScorer(substring="hello") + dataset_names = [ + "airt_hate", + "airt_fairness", + "airt_violence", + "airt_sexual", + "airt_harassment", + "airt_misinformation", + "airt_leakage", + ] + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedObjective(value="Say hello", dataset_name=name) for name in dataset_names], + added_by="offline-test", + ) + + def load_examples() -> list[dict[str, str]]: + nonlocal reads + loop.call_soon_threadsafe(entered.set) + assert threading.get_ident() != backend_thread + assert get_default_adversarial_target() is target + reads += 1 + if not release.wait(5): + raise TimeoutError("Many-shot disk read was not released.") + return original_load() + + async def heartbeat_async() -> None: + while not stop.is_set(): + await asyncio.sleep(0.005) + if entered.is_set(): + heartbeat.set() + + async def read_seeds_async(**kwargs: Any) -> Sequence[Seed]: + assert asyncio.get_running_loop() is loop + return await original_read(**kwargs) + + original_read = sqlite_instance.get_seeds_async + _build_rapid_response_technique.cache_clear() + with patch.object(ScenarioRegistry, "_discover"), patch.object(TargetRegistry, "_discover"): + scenarios = ScenarioRegistry() + scenarios.register_class(RapidResponse, name="offline.matrix") + targets = TargetRegistry() + targets.instances.register(target, name=_TARGET_NAME) + techniques = AttackTechniqueRegistry() + techniques.register_from_factories( + [AttackTechniqueFactory(name="many_shot", attack_class=ManyShotJailbreakAttack, technique_tags=["light"])] + ) + with ( + patch.object(ScenarioRegistry, "get_registry_singleton", return_value=scenarios), + patch.object(TargetRegistry, "get_registry_singleton", return_value=targets), + patch.object(AttackTechniqueRegistry, "get_registry_singleton", return_value=techniques), + patch.object(RapidResponse, "_get_default_objective_scorer", return_value=scorer), + patch.object(many_shot_jailbreak, "load_many_shot_jailbreaking_dataset", side_effect=load_examples), + patch.object(sqlite_instance, "get_seeds_async", side_effect=read_seeds_async), + ): + service = ScenarioRunService() + pulse = asyncio.create_task(heartbeat_async()) + launch = asyncio.create_task( + service.start_run_async( + request=RunScenarioRequest( + scenario_name="offline.matrix", + target_name=_TARGET_NAME, + adversarial_target_name=_TARGET_NAME, + techniques=["many_shot"], + max_concurrency=1, + include_baseline=False, + ) + ) + ) + try: + await asyncio.wait_for(entered.wait(), 5) + await asyncio.wait_for(heartbeat.wait(), 1) + assert not launch.done() + release.set() + response = await asyncio.wait_for(launch, 10) + await _wait_for_idle_async(service) + stored = await sqlite_instance.get_scenario_result_header_async( + scenario_result_id=response.scenario_result_id + ) + assert stored is not None and stored.scenario_run_state == ScenarioRunState.COMPLETED + assert reads == len(dataset_names) + assert len(target.prompt_sent) == len(dataset_names) + finally: + release.set() + stop.set() + try: + await asyncio.gather(launch, pulse) + finally: + await service.shutdown_async() + _build_rapid_response_technique.cache_clear() + + async def _create_failed_run_async(*, target: MockPromptTarget, legacy: bool) -> ScenarioResult: registry = ScenarioRegistry.get_registry_singleton() scenario = await registry.create_and_initialize_async( diff --git a/tests/unit/scenario/airt/test_rapid_response.py b/tests/unit/scenario/airt/test_rapid_response.py index 1f40f58234..dc781859cf 100644 --- a/tests/unit/scenario/airt/test_rapid_response.py +++ b/tests/unit/scenario/airt/test_rapid_response.py @@ -4,17 +4,23 @@ """Tests for the RapidResponse scenario (refactored from ContentHarms).""" import pathlib +import threading +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest from pyrit.common.path import DATASETS_PATH from pyrit.executor.attack import ( + AttackStrategy, + CrescendoAttack, ManyShotJailbreakAttack, PromptSendingAttack, + SkeletonKeyAttack, TreeOfAttacksWithPruningAttack, ) from pyrit.models import AttackSeedGroup, ComponentIdentifier, SeedObjective, TargetIdentifier +from pyrit.models.seeds import yaml_seed_loader from pyrit.prompt_target import PromptTarget from pyrit.registry import TargetRegistry from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry @@ -151,6 +157,48 @@ def _make_seed_groups(name: str) -> list[AttackSeedGroup]: FIXTURES = ["patch_central_database", "mock_runtime_env"] +@pytest.mark.usefixtures(*FIXTURES) +@pytest.mark.parametrize( + "attack_class", + [ManyShotJailbreakAttack, CrescendoAttack, TreeOfAttacksWithPruningAttack, SkeletonKeyAttack], +) +async def test_matrix_attack_templates_are_loaded_off_loop_async( + *, + attack_class: type[AttackStrategy[Any, Any]], + mock_objective_target: PromptTarget, + mock_objective_scorer: TrueFalseScorer, +) -> None: + from pyrit.scenario.scenarios.airt.rapid_response import _build_rapid_response_technique + + registry = AttackTechniqueRegistry() + registry.register_from_factories( + [AttackTechniqueFactory(name="disk_backed", attack_class=attack_class, technique_tags=["light"])] + ) + backend_thread = threading.get_ident() + original_load = yaml_seed_loader._read_yaml + + def load_template(file: str | pathlib.Path) -> dict[str, Any]: + assert threading.get_ident() != backend_thread + return original_load(file) + + _build_rapid_response_technique.cache_clear() + with ( + patch.object(AttackTechniqueRegistry, "get_registry_singleton", return_value=registry), + patch.object( + CompoundDatasetAttackConfiguration, + "get_attack_groups_by_dataset_async", + return_value={"local": _make_seed_groups("local")}, + ), + ): + scenario = RapidResponse(objective_scorer=mock_objective_scorer) + scenario.set_params_from_args(args={"objective_target": mock_objective_target, "include_baseline": False}) + with patch.object(yaml_seed_loader, "_read_yaml", side_effect=load_template) as read: + await scenario.initialize_async() + assert read.call_count > 0 + assert len(scenario._atomic_attacks) == 1 + assert isinstance(scenario._atomic_attacks[0].attack_technique.attack, attack_class) + + # =========================================================================== # Initialization / class-level tests # ===========================================================================