From 77825c623f30411bce624d4a47deeb4beca317b8 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 30 Sep 2026 16:48:29 -0400 Subject: [PATCH 01/16] FEAT: Add ScenarioPreset model, storage, and registry A scenario preset is a named, reusable, target-agnostic answer to *what to test*: a scenario plus its scenario-owned configuration. It deliberately excludes everything environment-specific (target, concurrency, retries, labels), which belongs to a launch. That split is what lets one preset run unchanged against dev, staging, and production. This is the first PR of the composite scan work and adds only the persistence layer. No API routes, no UI, and no built-in presets are registered yet. - `ScenarioPreset` / `ScenarioPresetProvenance` models. Every configurable field is tri-state: `None` means "not set by this preset, use the scenario default", which is distinct from an explicit value that happens to equal that default. Collapsing the two would pin a scenario default at save time and stop it tracking upstream changes. - `ScenarioPresetStorage` for JSON persistence with an optimistic-concurrency check on `version`. The caller states intent through a separate `expected_version` argument rather than through the version on the submitted model, so create and update are never ambiguous and a client-supplied version is never trusted into storage. - `ScenarioPresetRegistry` unioning built-in presets (registered from initializer code, read-only) with user presets from storage. Built-in wins a name collision and the colliding user preset is skipped with a warning, so shipping a new built-in cannot break a running install. - `ScenarioPresetConflictError` (HTTP 409) carrying both versions. - `FileDocumentStorage`, extracted from `CustomInitializerStorage`, holding the shared local-directory and Azure Blob handling for flat named documents. Extracted rather than duplicated because a second copy would mean two implementations of SAS detection, credential lifecycle, and container-URL parsing. `CustomInitializerStorage` now subclasses it with its public API unchanged, so its existing tests cover the refactor. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/exceptions/__init__.py | 2 + pyrit/exceptions/exception_classes.py | 33 ++ pyrit/models/__init__.py | 4 + pyrit/models/catalog/__init__.py | 3 + pyrit/models/catalog/scenario_preset.py | 92 ++++++ pyrit/registry/__init__.py | 4 + pyrit/registry/custom_initializer_storage.py | 136 +-------- pyrit/registry/file_document_storage.py | 214 +++++++++++++ pyrit/registry/scenario_preset_registry.py | 177 +++++++++++ pyrit/registry/scenario_preset_storage.py | 169 +++++++++++ tests/unit/models/test_scenario_preset.py | 91 ++++++ .../test_custom_initializer_storage.py | 16 +- .../registry/test_scenario_preset_registry.py | 206 +++++++++++++ .../registry/test_scenario_preset_storage.py | 286 ++++++++++++++++++ 14 files changed, 1300 insertions(+), 133 deletions(-) create mode 100644 pyrit/models/catalog/scenario_preset.py create mode 100644 pyrit/registry/file_document_storage.py create mode 100644 pyrit/registry/scenario_preset_registry.py create mode 100644 pyrit/registry/scenario_preset_storage.py create mode 100644 tests/unit/models/test_scenario_preset.py create mode 100644 tests/unit/registry/test_scenario_preset_registry.py create mode 100644 tests/unit/registry/test_scenario_preset_storage.py diff --git a/pyrit/exceptions/__init__.py b/pyrit/exceptions/__init__.py index a42582594b..dbae736cbe 100644 --- a/pyrit/exceptions/__init__.py +++ b/pyrit/exceptions/__init__.py @@ -22,6 +22,7 @@ PyritException, RateLimitException, ScenarioPartialFailureException, + ScenarioPresetConflictError, ScorerLLMResponseBlockedException, get_retry_max_num_attempts, handle_bad_request_exception, @@ -77,6 +78,7 @@ "remove_markdown_json": "pyrit.exceptions.exceptions_helpers", "RetryCollector": "pyrit.exceptions.retry_collector", "ScenarioPartialFailureException": "pyrit.exceptions.exception_classes", + "ScenarioPresetConflictError": "pyrit.exceptions.exception_classes", "ScorerLLMResponseBlockedException": "pyrit.exceptions.exception_classes", "set_execution_context": "pyrit.exceptions.exception_context", "set_retry_collector": "pyrit.exceptions.retry_collector", diff --git a/pyrit/exceptions/exception_classes.py b/pyrit/exceptions/exception_classes.py index b15fa8272c..483fe4f70e 100644 --- a/pyrit/exceptions/exception_classes.py +++ b/pyrit/exceptions/exception_classes.py @@ -302,6 +302,39 @@ def __init__( self.__cause__ = self.incomplete_objectives[0][1] +class ScenarioPresetConflictError(PyritException): + """ + Exception raised when a scenario preset save loses an optimistic-concurrency check. + + Carries the version the caller edited against and the version currently stored so a + caller can show the user what changed underneath them rather than silently overwriting. + """ + + def __init__(self, *, name: str, expected_version: int | None, actual_version: int | None) -> None: + """ + Initialize a scenario preset conflict error. + + Args: + name (str): Name of the preset that could not be saved. + expected_version (int | None): Version the caller based its edit on, or ``None`` + when the caller intended to create a new preset. + actual_version (int | None): Version currently stored, or ``None`` when no preset + with that name exists. + """ + self.name = name + self.expected_version = expected_version + self.actual_version = actual_version + + if actual_version is None: + detail = f"expected version {expected_version} but it no longer exists" + elif expected_version is None: + detail = f"it already exists at version {actual_version}" + else: + detail = f"expected version {expected_version} but found {actual_version}" + + super().__init__(status_code=409, message=f"Scenario preset '{name}' could not be saved: {detail}.") + + class InvalidJsonException(PyritException): """Exception class for blocked content errors.""" diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index bc34c01116..f5342a6996 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -49,6 +49,8 @@ ScenarioDatasetSizeCap, ScenarioDatasetSummary, ScenarioDefaultRunSizeEstimate, + ScenarioPreset, + ScenarioPresetProvenance, ScenarioRunListItem, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, @@ -356,6 +358,8 @@ "ScenarioDatasetSizeCap": "pyrit.models.catalog", "ScenarioDatasetSummary": "pyrit.models.catalog", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog", + "ScenarioPreset": "pyrit.models.catalog", + "ScenarioPresetProvenance": "pyrit.models.catalog", "ScenarioRunListItem": "pyrit.models.catalog", "ScenarioRunSizeComponent": "pyrit.models.catalog", "ScenarioRunSizeEstimate": "pyrit.models.catalog", diff --git a/pyrit/models/catalog/__init__.py b/pyrit/models/catalog/__init__.py index c7a807f039..f0739608f5 100644 --- a/pyrit/models/catalog/__init__.py +++ b/pyrit/models/catalog/__init__.py @@ -38,6 +38,7 @@ ScenarioRunSummary, ScenarioTechniqueSummary, ) + from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance from pyrit.models.catalog.target import TargetInstance _LAZY_EXPORTS: dict[str, str] = { @@ -49,6 +50,8 @@ "ScenarioDatasetSizeCap": "pyrit.models.catalog.scenario", "ScenarioDatasetSummary": "pyrit.models.catalog.scenario", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog.scenario", + "ScenarioPreset": "pyrit.models.catalog.scenario_preset", + "ScenarioPresetProvenance": "pyrit.models.catalog.scenario_preset", "ScenarioRunListItem": "pyrit.models.catalog.scenario", "ScenarioRunSizeComponent": "pyrit.models.catalog.scenario", "ScenarioRunSizeEstimate": "pyrit.models.catalog.scenario", diff --git a/pyrit/models/catalog/scenario_preset.py b/pyrit/models/catalog/scenario_preset.py new file mode 100644 index 0000000000..8ac955b470 --- /dev/null +++ b/pyrit/models/catalog/scenario_preset.py @@ -0,0 +1,92 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Scenario preset model. + +A scenario preset is a named, reusable, target-agnostic answer to *what to test*: +a scenario plus the scenario-owned configuration fields. It deliberately excludes +everything environment-specific — target, concurrency, retries, and labels — which +belongs to a launch rather than to the preset. That split is what lets one preset +run unchanged against dev, staging, and production. + +Every configurable field is tri-state. ``None`` means "not set by this preset, use +the scenario's own default", which is distinct from an explicit value that happens +to equal that default. Collapsing the two would silently pin a scenario default at +the moment a preset was saved and stop it tracking upstream changes. +""" + +from enum import Enum +from typing import Any + +from pydantic import BaseModel, Field, field_validator + +from pyrit.models.catalog.scenario import _validate_dataset_filter_mapping +from pyrit.models.identifiers.class_name_utils import validate_registry_name + + +class ScenarioPresetProvenance(str, Enum): + """Where a preset came from, which determines whether it can be edited.""" + + BUILT_IN = "built_in" + USER = "user" + + +class ScenarioPreset(BaseModel): + """ + A named, reusable, target-agnostic scenario configuration. + + Presets own *what to test*. A launch owns *how and where* — the target, + concurrency, retries, and labels — so those fields are absent here by design + and the two sets are combined by union rather than by precedence. + """ + + name: str = Field(..., description="Unique preset name, used as the storage key and as the reference from scans") + scenario_name: str = Field(..., min_length=1, description="Registered scenario this preset configures") + description: str | None = Field(None, description="Human-readable summary of what this preset tests") + techniques: list[str] | None = Field(None, description="Technique names; None uses the scenario default") + dataset_names: list[str] | None = Field(None, description="Dataset names; None uses the scenario default") + max_dataset_size: int | None = Field(None, ge=1, description="Maximum selected logical seed groups") + dataset_filters: dict[str, list[str]] | None = Field( + None, + description="Dataset seed filters keyed by field. Accepted keys: harm_categories, data_types.", + ) + include_baseline: bool | None = Field(None, description="Override the scenario baseline default") + scenario_params: dict[str, Any] | None = Field( + None, description="Scenario-declared parameters such as template names and attempt counts" + ) + version: int = Field(1, ge=1, description="Monotonic change counter assigned by storage, not an address") + provenance: ScenarioPresetProvenance = Field( + ScenarioPresetProvenance.USER, description="Whether this preset ships with PyRIT or was authored by a user" + ) + + @property + def is_builtin(self) -> bool: + """Whether this preset ships with PyRIT and is therefore read-only.""" + return self.provenance is ScenarioPresetProvenance.BUILT_IN + + @field_validator("name") + @classmethod + def _validate_name(cls, value: str) -> str: + """ + Validate that the preset name is a legal registry name. + + Returns: + str: The validated name. + + Raises: + ValueError: If the name is not a legal registry name. + """ + validate_registry_name(value) + return value + + @field_validator("dataset_filters") + @classmethod + def _validate_dataset_filters(cls, value: dict[str, list[str]] | None) -> dict[str, list[str]] | None: + """ + Validate dataset filters against the shared allow-list. + + Returns: + dict[str, list[str]] | None: Validated filters. + """ + return _validate_dataset_filter_mapping(value) diff --git a/pyrit/registry/__init__.py b/pyrit/registry/__init__.py index ed782684e7..d6bb1afb58 100644 --- a/pyrit/registry/__init__.py +++ b/pyrit/registry/__init__.py @@ -32,6 +32,8 @@ ) from pyrit.registry.registry import InstanceHoldingRegistry, ParamBagRegistry, Registry from pyrit.registry.registry_metadata import RegistryMetadata + from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry + from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage from pyrit.registry.tag_query import TagQuery _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { @@ -51,6 +53,8 @@ "InitializerRegistry": "pyrit.registry.components", "RegistryEntry": "pyrit.registry.instance_registry", "ScenarioMetadata": "pyrit.registry.components", + "ScenarioPresetRegistry": "pyrit.registry.scenario_preset_registry", + "ScenarioPresetStorage": "pyrit.registry.scenario_preset_storage", "ScenarioRegistry": "pyrit.registry.components", "ScorerRegistry": "pyrit.registry.components", "ScorerMetadata": "pyrit.registry.components", diff --git a/pyrit/registry/custom_initializer_storage.py b/pyrit/registry/custom_initializer_storage.py index 90e63ed7a0..8b38323312 100644 --- a/pyrit/registry/custom_initializer_storage.py +++ b/pyrit/registry/custom_initializer_storage.py @@ -5,20 +5,10 @@ from __future__ import annotations -from contextlib import contextmanager, suppress -from pathlib import Path, PurePosixPath -from typing import TYPE_CHECKING -from urllib.parse import unquote, urlparse +from pyrit.registry.file_document_storage import FileDocumentStorage -from pyrit.common.azure_storage import has_sas_signature, is_azure_blob_uri, redact_url_credentials -if TYPE_CHECKING: - from collections.abc import Generator - - from azure.storage.blob import ContainerClient - - -class CustomInitializerStorage: +class CustomInitializerStorage(FileDocumentStorage): """Read and write custom initializer scripts in a directory or blob container.""" def __init__(self, *, source: str) -> None: @@ -28,21 +18,7 @@ def __init__(self, *, source: str) -> None: Raises: ValueError: If the source has an unsupported URI scheme. """ - self._source = source - self._is_blob = is_azure_blob_uri(source) - if not self._is_blob and urlparse(source).scheme and not Path(source).drive: - raise ValueError( - "Custom initializer source must be a local directory or Azure Blob container URI " - "with an optional blob prefix" - ) - self._container_url, self._blob_prefix = self._parse_blob_source() if self._is_blob else (None, "") - - @property - def display_source(self) -> str: - """Storage source without Azure Blob credentials.""" - if not self._is_blob: - return self._source - return redact_url_credentials(self._source) + super().__init__(source=source, extension=".py", source_label="Custom initializer") def get_script_source(self, name: str) -> str: """ @@ -51,9 +27,7 @@ def get_script_source(self, name: str) -> str: Returns: str: Local file path or Azure Blob URI for the script. """ - if self._is_blob: - return f"{self.display_source.rstrip('/')}/{name}.py" - return str(Path(self._source).expanduser() / f"{name}.py") + return self._get_document_source(name) def list_scripts(self) -> dict[str, str]: """ @@ -62,108 +36,12 @@ def list_scripts(self) -> dict[str, str]: Returns: dict[str, str]: Script content keyed by registry name. """ - if self._is_blob: - return self._list_blob_scripts() - - directory = Path(self._source).expanduser() - directory.mkdir(parents=True, exist_ok=True) - return {path.stem: path.read_text(encoding="utf-8") for path in sorted(directory.glob("*.py"))} + return self._list_documents() def save_script(self, *, name: str, content: str) -> None: """Persist one custom initializer script.""" - if self._is_blob: - with self._open_container_client() as client: - client.upload_blob(name=self._get_blob_name(name), data=content.encode("utf-8"), overwrite=True) - else: - directory = Path(self._source).expanduser() - directory.mkdir(parents=True, exist_ok=True) - (directory / f"{name}.py").write_text(content, encoding="utf-8") + self._save_document(name=name, content=content) def delete_script(self, name: str) -> None: """Delete one custom initializer script if it exists.""" - if self._is_blob: - from azure.core.exceptions import ResourceNotFoundError - - with self._open_container_client() as client: - with suppress(ResourceNotFoundError): - client.delete_blob(self._get_blob_name(name)) - else: - (Path(self._source).expanduser() / f"{name}.py").unlink(missing_ok=True) - - def _list_blob_scripts(self) -> dict[str, str]: - """ - Read Python scripts from the configured Azure Blob container. - - Returns: - dict[str, str]: Script content keyed by blob stem. - """ - scripts: dict[str, str] = {} - with self._open_container_client() as client: - prefix = f"{self._blob_prefix}/" if self._blob_prefix else None - blobs = client.list_blobs(name_starts_with=prefix) if prefix else client.list_blobs() - blob_names = sorted( - blob.name for blob in blobs if self._is_direct_python_blob(blob_name=blob.name, prefix=prefix) - ) - for blob_name in blob_names: - relative_name = blob_name.removeprefix(prefix or "") - scripts[PurePosixPath(relative_name).stem] = client.download_blob(blob_name).readall().decode("utf-8") - return scripts - - def _parse_blob_source(self) -> tuple[str, str]: - """ - Split the configured source into a container URL and blob prefix. - - Returns: - tuple[str, str]: The container URL and decoded blob prefix. - """ - parsed_uri = urlparse(self._source) - container_path, _, prefix = parsed_uri.path.strip("/").partition("/") - container_url = parsed_uri._replace(path=f"/{container_path}", fragment="").geturl() - return container_url, unquote(prefix).strip("/") - - def _get_blob_name(self, name: str) -> str: - """ - Build the blob name for a registry entry. - - Returns: - str: The prefixed Python blob name. - """ - file_name = f"{name}.py" - return f"{self._blob_prefix}/{file_name}" if self._blob_prefix else file_name - - @staticmethod - def _is_direct_python_blob(*, blob_name: str, prefix: str | None) -> bool: - """Return whether a blob is a direct Python child of the configured prefix.""" - if prefix and not blob_name.startswith(prefix): - return False - relative_name = blob_name.removeprefix(prefix or "") - return "/" not in relative_name and PurePosixPath(relative_name).suffix == ".py" - - @contextmanager - def _open_container_client(self) -> Generator[ContainerClient, None, None]: - """ - Yield an Azure Blob container client and close its credential. - - Yields: - ContainerClient: A client scoped to the configured container. - - Raises: - RuntimeError: If called for a non-Blob source. - """ - from azure.identity import DefaultAzureCredential - from azure.storage.blob import ContainerClient - - if self._container_url is None: - raise RuntimeError("Azure Blob container URL is not configured") - - if has_sas_signature(self._container_url): - with ContainerClient.from_container_url(container_url=self._container_url) as client: - yield client - return - - with DefaultAzureCredential() as credential: - with ContainerClient.from_container_url( - container_url=self._container_url, - credential=credential, - ) as client: - yield client + self._delete_document(name) diff --git a/pyrit/registry/file_document_storage.py b/pyrit/registry/file_document_storage.py new file mode 100644 index 0000000000..25cade9d78 --- /dev/null +++ b/pyrit/registry/file_document_storage.py @@ -0,0 +1,214 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Shared local-directory and Azure Blob storage for flat, named documents.""" + +from __future__ import annotations + +from contextlib import contextmanager, suppress +from pathlib import Path, PurePosixPath +from typing import TYPE_CHECKING +from urllib.parse import unquote, urlparse + +from pyrit.common.azure_storage import has_sas_signature, is_azure_blob_uri, redact_url_credentials + +if TYPE_CHECKING: + from collections.abc import Generator + + from azure.storage.blob import ContainerClient + + +class FileDocumentStorage: + """ + Read and write flat, named text documents in a directory or blob container. + + Documents live directly under the configured source, one file per name, with a + fixed extension. Nested paths are ignored so a virtual directory prefix behaves + the same way in both backends. + + Subclasses supply the extension and a human-readable label for error messages, + then expose a domain-specific API over the protected document operations. + """ + + def __init__(self, *, source: str, extension: str, source_label: str) -> None: + """ + Initialize storage from a local directory or Azure Blob source URI. + + Args: + source (str): Local directory path or Azure Blob container URI with an + optional blob prefix. + extension (str): File extension including the leading dot, such as ``".py"``. + source_label (str): Human-readable label naming the stored documents, used + in error messages. + + Raises: + ValueError: If the source has an unsupported URI scheme. + """ + self._source = source + self._extension = extension + self._is_blob = is_azure_blob_uri(source) + if not self._is_blob and urlparse(source).scheme and not Path(source).drive: + raise ValueError( + f"{source_label} source must be a local directory or Azure Blob container URI " + "with an optional blob prefix" + ) + self._container_url, self._blob_prefix = self._parse_blob_source() if self._is_blob else (None, "") + + @property + def display_source(self) -> str: + """Storage source without Azure Blob credentials.""" + if not self._is_blob: + return self._source + return redact_url_credentials(self._source) + + def _get_document_source(self, name: str) -> str: + """ + Get the credential-free location of one document. + + Returns: + str: Local file path or Azure Blob URI for the document. + """ + if self._is_blob: + return f"{self.display_source.rstrip('/')}/{name}{self._extension}" + return str(self._local_directory() / f"{name}{self._extension}") + + def _list_documents(self) -> dict[str, str]: + """ + Read every stored document. + + Returns: + dict[str, str]: Document content keyed by name. + """ + if self._is_blob: + return self._list_blob_documents() + + directory = self._local_directory(create=True) + return {path.stem: path.read_text(encoding="utf-8") for path in sorted(directory.glob(f"*{self._extension}"))} + + def _read_document(self, name: str) -> str | None: + """ + Read one document. + + Returns: + str | None: Document content, or ``None`` if it does not exist. + """ + if self._is_blob: + from azure.core.exceptions import ResourceNotFoundError + + with self._open_container_client() as client: + try: + return client.download_blob(self._get_blob_name(name)).readall().decode("utf-8") + except ResourceNotFoundError: + return None + + path = self._local_directory() / f"{name}{self._extension}" + return path.read_text(encoding="utf-8") if path.is_file() else None + + def _save_document(self, *, name: str, content: str) -> None: + """Persist one document, overwriting any existing content.""" + if self._is_blob: + with self._open_container_client() as client: + client.upload_blob(name=self._get_blob_name(name), data=content.encode("utf-8"), overwrite=True) + else: + directory = self._local_directory(create=True) + (directory / f"{name}{self._extension}").write_text(content, encoding="utf-8") + + def _delete_document(self, name: str) -> None: + """Delete one document if it exists.""" + if self._is_blob: + from azure.core.exceptions import ResourceNotFoundError + + with self._open_container_client() as client: + with suppress(ResourceNotFoundError): + client.delete_blob(self._get_blob_name(name)) + else: + (self._local_directory() / f"{name}{self._extension}").unlink(missing_ok=True) + + def _local_directory(self, *, create: bool = False) -> Path: + """ + Resolve the configured local directory. + + Returns: + Path: The expanded directory path. + """ + directory = Path(self._source).expanduser() + if create: + directory.mkdir(parents=True, exist_ok=True) + return directory + + def _list_blob_documents(self) -> dict[str, str]: + """ + Read documents from the configured Azure Blob container. + + Returns: + dict[str, str]: Document content keyed by blob stem. + """ + documents: dict[str, str] = {} + with self._open_container_client() as client: + prefix = f"{self._blob_prefix}/" if self._blob_prefix else None + blobs = client.list_blobs(name_starts_with=prefix) if prefix else client.list_blobs() + blob_names = sorted( + blob.name for blob in blobs if self._is_direct_document_blob(blob_name=blob.name, prefix=prefix) + ) + for blob_name in blob_names: + relative_name = blob_name.removeprefix(prefix or "") + documents[PurePosixPath(relative_name).stem] = client.download_blob(blob_name).readall().decode("utf-8") + return documents + + def _parse_blob_source(self) -> tuple[str, str]: + """ + Split the configured source into a container URL and blob prefix. + + Returns: + tuple[str, str]: The container URL and decoded blob prefix. + """ + parsed_uri = urlparse(self._source) + container_path, _, prefix = parsed_uri.path.strip("/").partition("/") + container_url = parsed_uri._replace(path=f"/{container_path}", fragment="").geturl() + return container_url, unquote(prefix).strip("/") + + def _get_blob_name(self, name: str) -> str: + """ + Build the blob name for one document. + + Returns: + str: The prefixed blob name. + """ + file_name = f"{name}{self._extension}" + return f"{self._blob_prefix}/{file_name}" if self._blob_prefix else file_name + + def _is_direct_document_blob(self, *, blob_name: str, prefix: str | None) -> bool: + """Return whether a blob is a direct child of the configured prefix with the expected extension.""" + if prefix and not blob_name.startswith(prefix): + return False + relative_name = blob_name.removeprefix(prefix or "") + return "/" not in relative_name and PurePosixPath(relative_name).suffix == self._extension + + @contextmanager + def _open_container_client(self) -> Generator[ContainerClient, None, None]: + """ + Yield an Azure Blob container client and close its credential. + + Yields: + ContainerClient: A client scoped to the configured container. + + Raises: + RuntimeError: If called for a non-Blob source. + """ + from azure.identity import DefaultAzureCredential + from azure.storage.blob import ContainerClient + + if self._container_url is None: + raise RuntimeError("Azure Blob container URL is not configured") + + if has_sas_signature(self._container_url): + with ContainerClient.from_container_url(container_url=self._container_url) as client: + yield client + return + + with DefaultAzureCredential() as credential: + with ContainerClient.from_container_url( + container_url=self._container_url, + credential=credential, + ) as client: + yield client diff --git a/pyrit/registry/scenario_preset_registry.py b/pyrit/registry/scenario_preset_registry.py new file mode 100644 index 0000000000..1ff2c8c82e --- /dev/null +++ b/pyrit/registry/scenario_preset_registry.py @@ -0,0 +1,177 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""In-memory registry unioning built-in and user-authored scenario presets.""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage + +if TYPE_CHECKING: + from pathlib import Path + +logger = logging.getLogger(__name__) + + +class ScenarioPresetRegistry: + """ + The union of built-in and user-authored scenario presets. + + Presets arrive from two places. Built-in presets are registered from initializer code + at startup and are read-only. User presets are loaded from storage. The registry exists + because only it sees both, so only it can resolve a name or detect a dangling reference. + + Built-in presets win on a name collision. A user preset that collides is skipped with a + warning rather than failing startup, because the collision arises when PyRIT ships a new + built-in under a name a user already took, and an upgrade must not break a running + install. To customize a built-in, fork it under a new name. + """ + + def __init__(self, *, storage: ScenarioPresetStorage | None = None) -> None: + """ + Initialize the registry. + + Args: + storage (ScenarioPresetStorage | None): Storage for user presets. Defaults to a + local directory under the PyRIT configuration path. + """ + self._storage = storage + self._builtin_presets: dict[str, ScenarioPreset] = {} + self._user_presets: dict[str, ScenarioPreset] = {} + + def configure_storage_source(self, source: str | None) -> None: + """Configure the local directory or Azure Blob source for user presets.""" + self._storage = ScenarioPresetStorage(source=source or str(self._get_default_storage_dir())) + + def register_builtin(self, preset: ScenarioPreset) -> ScenarioPreset: + """ + Register a preset that ships with PyRIT. + + Args: + preset (ScenarioPreset): The preset to register. Its provenance is forced to + built-in so a caller cannot register an editable preset by accident. + + Returns: + ScenarioPreset: The registered preset. + + Raises: + ValueError: If a built-in preset with the same name is already registered. + """ + if preset.name in self._builtin_presets: + raise ValueError(f"Built-in scenario preset '{preset.name}' is already registered.") + + registered = preset.model_copy(update={"provenance": ScenarioPresetProvenance.BUILT_IN}) + self._builtin_presets[registered.name] = registered + return registered + + def load_stored_presets(self) -> None: + """Load user presets from storage, skipping any that collide with a built-in preset.""" + self._user_presets = {} + for name, preset in self._get_storage().list_presets().items(): + if name in self._builtin_presets: + logger.warning( + f"Skipping stored scenario preset '{name}': a built-in preset already uses that name. " + "Rename the stored preset to keep using it." + ) + continue + self._user_presets[name] = preset + + def get_preset(self, name: str) -> ScenarioPreset | None: + """ + Resolve one preset by name. + + Returns: + ScenarioPreset | None: The preset, or ``None`` if no preset uses that name. + """ + return self._builtin_presets.get(name) or self._user_presets.get(name) + + def list_presets(self) -> list[ScenarioPreset]: + """ + List every known preset. + + Returns: + list[ScenarioPreset]: Built-in and user presets, sorted by name. + """ + merged = {**self._user_presets, **self._builtin_presets} + return [merged[name] for name in sorted(merged)] + + def is_builtin(self, name: str) -> bool: + """Return whether *name* belongs to a built-in preset.""" + return name in self._builtin_presets + + def save_preset(self, *, preset: ScenarioPreset, expected_version: int | None) -> ScenarioPreset: + """ + Persist a user preset and refresh the in-memory copy. + + Args: + preset (ScenarioPreset): The preset to persist. + expected_version (int | None): ``None`` to create a preset that must not already + exist, or the version the edit was based on. + + Returns: + ScenarioPreset: The persisted preset, carrying its newly assigned version. + + Raises: + ScenarioPresetConflictError: If the stored version does not match *expected_version*. + ValueError: If *name* belongs to a built-in preset. + """ + self._reject_builtin(preset.name, action="saved") + + saved = self._get_storage().save_preset(preset=preset, expected_version=expected_version) + self._user_presets[saved.name] = saved + return saved + + def delete_preset(self, name: str) -> None: + """ + Delete a user preset from storage and from the registry. + + Raises: + ValueError: If *name* belongs to a built-in preset. + """ + self._reject_builtin(name, action="deleted") + + self._get_storage().delete_preset(name) + self._user_presets.pop(name, None) + + def _reject_builtin(self, name: str, *, action: str) -> None: + """ + Refuse to mutate a built-in preset. + + Raises: + ValueError: If *name* belongs to a built-in preset. + """ + if name in self._builtin_presets: + raise ValueError( + f"Scenario preset '{name}' is built in and cannot be {action}. " + "Fork it under a new name to customize it." + ) + + def _get_storage(self) -> ScenarioPresetStorage: + """ + Return storage for user presets, creating the default backend on first use. + + Returns: + ScenarioPresetStorage: The configured storage backend. + """ + if self._storage is None: + self._storage = ScenarioPresetStorage(source=str(self._get_default_storage_dir())) + return self._storage + + @staticmethod + def _get_default_storage_dir() -> Path: + """ + Get the directory for storing user-authored presets. + + Returns: + Path: Path to ``~/.pyrit/scenario_presets/``, created if needed. + """ + # Deferred: importing pyrit.common.path triggers pyrit __init__.py + from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH + + presets_dir = CONFIGURATION_DIRECTORY_PATH / "scenario_presets" + presets_dir.mkdir(parents=True, exist_ok=True) + return presets_dir diff --git a/pyrit/registry/scenario_preset_storage.py b/pyrit/registry/scenario_preset_storage.py new file mode 100644 index 0000000000..59984945c5 --- /dev/null +++ b/pyrit/registry/scenario_preset_storage.py @@ -0,0 +1,169 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Storage backends for user-authored scenario presets.""" + +from __future__ import annotations + +import json +import logging + +from pyrit.exceptions.exception_classes import ScenarioPresetConflictError +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.registry.file_document_storage import FileDocumentStorage + +logger = logging.getLogger(__name__) + + +class ScenarioPresetStorage(FileDocumentStorage): + """ + Read and write user-authored scenario presets as JSON documents. + + Only user presets are stored. Built-in presets come from initializer code and are + never written here, so a file on disk always means a user authored it. + + Writes are guarded by an optimistic-concurrency check on ``version``. The stored + version is authoritative: a caller supplies the version it edited against, and the + save is refused if storage has moved on. + + The check is read-then-write rather than a true compare-and-swap, so two saves racing + within the same instant can both observe the same version and the later write wins. + It is aimed at the realistic case — a person editing a copy that went stale minutes + ago — not at concurrent writers. Closing that gap needs backend-specific conditional + writes (blob ETags have no local-filesystem equivalent) and is deliberately deferred. + """ + + def __init__(self, *, source: str) -> None: + """ + Initialize storage from a local directory or Azure Blob source URI. + + Raises: + ValueError: If the source has an unsupported URI scheme. + """ + super().__init__(source=source, extension=".json", source_label="Scenario preset") + + def get_preset_source(self, name: str) -> str: + """ + Get the credential-free location of one stored preset. + + Returns: + str: Local file path or Azure Blob URI for the preset. + """ + return self._get_document_source(name) + + def list_presets(self) -> dict[str, ScenarioPreset]: + """ + Read every stored preset, skipping any that cannot be parsed. + + A malformed or hand-edited file must not prevent the rest of the library from + loading, so failures are logged and that preset is omitted. + + Returns: + dict[str, ScenarioPreset]: Presets keyed by name. + """ + presets: dict[str, ScenarioPreset] = {} + for name, content in self._list_documents().items(): + preset = self._parse_preset(name=name, content=content) + if preset is not None: + presets[preset.name] = preset + return presets + + def load_preset(self, name: str) -> ScenarioPreset | None: + """ + Read one stored preset. + + Returns: + ScenarioPreset | None: The preset, or ``None`` if it is absent or malformed. + """ + content = self._read_document(name) + if content is None: + return None + return self._parse_preset(name=name, content=content) + + def save_preset(self, *, preset: ScenarioPreset, expected_version: int | None) -> ScenarioPreset: + """ + Persist one preset, assigning its version. + + The caller states its intent through *expected_version* rather than through the + version on *preset*, so creating and updating are never ambiguous and a + client-supplied version can never be trusted into storage. + + Args: + preset (ScenarioPreset): The preset to persist. + expected_version (int | None): ``None`` to create a preset that must not already + exist, or the version the edit was based on. + + Returns: + ScenarioPreset: The persisted preset, carrying its newly assigned version. + + Raises: + ScenarioPresetConflictError: If the stored version does not match *expected_version*. + ValueError: If *preset* is built-in, which is never stored. + """ + if preset.is_builtin: + raise ValueError( + f"Scenario preset '{preset.name}' is built in and cannot be saved. " + "Fork it under a new name to customize it." + ) + + existing = self.load_preset(preset.name) + actual_version = existing.version if existing is not None else None + if actual_version != expected_version: + raise ScenarioPresetConflictError( + name=preset.name, expected_version=expected_version, actual_version=actual_version + ) + + saved = preset.model_copy( + update={ + "version": 1 if existing is None else existing.version + 1, + "provenance": ScenarioPresetProvenance.USER, + } + ) + self._save_document(name=saved.name, content=self._serialize_preset(saved)) + return saved + + def delete_preset(self, name: str) -> None: + """Delete one stored preset if it exists.""" + self._delete_document(name) + + @staticmethod + def _serialize_preset(preset: ScenarioPreset) -> str: + """ + Serialize a preset to stored JSON. + + Unset fields are omitted rather than written as ``null`` so a stored preset reads + as the set of decisions its author actually made. + + Returns: + str: JSON document content. + """ + return json.dumps(preset.model_dump(mode="json", exclude_none=True), indent=2, sort_keys=True) + "\n" + + @staticmethod + def _parse_preset(*, name: str, content: str) -> ScenarioPreset | None: + """ + Parse one stored preset document. + + The document name is authoritative, overriding any ``name`` inside the payload. + It is the storage key, so letting the payload disagree would make a load return a + preset whose subsequent save wrote to a different file. + + Returns: + ScenarioPreset | None: The parsed preset, or ``None`` if it is malformed. + """ + try: + payload = json.loads(content) + except ValueError: + logger.exception(f"Skipping stored scenario preset '{name}': it is not valid JSON.") + return None + + if not isinstance(payload, dict): + logger.error(f"Skipping stored scenario preset '{name}': it is not a JSON object.") + return None + + fields = {**payload, "name": name, "provenance": ScenarioPresetProvenance.USER} + try: + return ScenarioPreset.model_validate(fields) + except Exception: + logger.exception(f"Skipping stored scenario preset '{name}': it is not a valid preset.") + return None diff --git a/tests/unit/models/test_scenario_preset.py b/tests/unit/models/test_scenario_preset.py new file mode 100644 index 0000000000..480ce460c0 --- /dev/null +++ b/tests/unit/models/test_scenario_preset.py @@ -0,0 +1,91 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the scenario preset model.""" + +import pytest + +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance + + +def test_init_defaults_every_optional_field_to_none() -> None: + """Test that an unset preset field is None rather than a concrete default.""" + preset = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent") + + assert preset.techniques is None + assert preset.dataset_names is None + assert preset.max_dataset_size is None + assert preset.dataset_filters is None + assert preset.include_baseline is None + assert preset.scenario_params is None + assert preset.description is None + + +def test_init_defaults_to_user_provenance_at_version_one() -> None: + """Test the default identity of a freshly constructed preset.""" + preset = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent") + + assert preset.version == 1 + assert preset.provenance is ScenarioPresetProvenance.USER + assert preset.is_builtin is False + + +def test_is_builtin_reflects_provenance() -> None: + """Test that built-in provenance marks a preset read-only.""" + preset = ScenarioPreset( + name="nightly", + scenario_name="foundry.red_team_agent", + provenance=ScenarioPresetProvenance.BUILT_IN, + ) + + assert preset.is_builtin is True + + +@pytest.mark.parametrize("name", ["Nightly", "nightly-scan", "1nightly", "", "a" * 65]) +def test_init_rejects_invalid_registry_names(name: str) -> None: + """Test that a preset name must be a legal registry name.""" + with pytest.raises(ValueError, match="Invalid registry name"): + ScenarioPreset(name=name, scenario_name="foundry.red_team_agent") + + +def test_init_rejects_unknown_dataset_filter() -> None: + """Test that dataset filters are validated against the shared allow-list.""" + with pytest.raises(ValueError, match="Unknown dataset filter"): + ScenarioPreset( + name="nightly", + scenario_name="foundry.red_team_agent", + dataset_filters={"not_a_field": ["x"]}, + ) + + +def test_init_accepts_known_dataset_filters() -> None: + """Test that allow-listed dataset filters are preserved.""" + preset = ScenarioPreset( + name="nightly", + scenario_name="foundry.red_team_agent", + dataset_filters={"harm_categories": ["violence"], "data_types": ["text"]}, + ) + + assert preset.dataset_filters == {"harm_categories": ["violence"], "data_types": ["text"]} + + +def test_init_rejects_zero_max_dataset_size() -> None: + """Test that a dataset cap must select at least one item.""" + with pytest.raises(ValueError): + ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent", max_dataset_size=0) + + +def test_init_rejects_empty_scenario_name() -> None: + """Test that a preset must name the scenario it configures.""" + with pytest.raises(ValueError): + ScenarioPreset(name="nightly", scenario_name="") + + +def test_include_baseline_distinguishes_unset_from_false() -> None: + """Test the tri-state contract that an unset override is not a disabled override.""" + unset = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent") + disabled = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent", include_baseline=False) + + assert unset.include_baseline is None + assert disabled.include_baseline is False + assert unset.include_baseline != disabled.include_baseline diff --git a/tests/unit/registry/test_custom_initializer_storage.py b/tests/unit/registry/test_custom_initializer_storage.py index f739375ff0..cace9d16a2 100644 --- a/tests/unit/registry/test_custom_initializer_storage.py +++ b/tests/unit/registry/test_custom_initializer_storage.py @@ -135,11 +135,19 @@ def test_local_storage_reads_latest_script_content(tmp_path: Path) -> None: assert storage.list_scripts() == {"example": "VALUE = 2\n"} -def test_direct_python_blob_rejects_name_outside_prefix() -> None: +def test_direct_python_blob_rejects_name_outside_prefix(tmp_path: Path) -> None: """Test that blobs outside the configured virtual directory are ignored.""" - assert not CustomInitializerStorage._is_direct_python_blob( - blob_name="other/example.py", prefix="custom-initializers/" - ) + storage = CustomInitializerStorage(source=str(tmp_path)) + + assert not storage._is_direct_document_blob(blob_name="other/example.py", prefix="custom-initializers/") + + +def test_direct_document_blob_rejects_other_extensions(tmp_path: Path) -> None: + """Test that only blobs with the configured extension are listed.""" + storage = CustomInitializerStorage(source=str(tmp_path)) + + assert storage._is_direct_document_blob(blob_name="example.py", prefix=None) + assert not storage._is_direct_document_blob(blob_name="example.json", prefix=None) def test_local_storage_cannot_open_blob_client(tmp_path: Path) -> None: diff --git a/tests/unit/registry/test_scenario_preset_registry.py b/tests/unit/registry/test_scenario_preset_registry.py new file mode 100644 index 0000000000..6a8719d2fc --- /dev/null +++ b/tests/unit/registry/test_scenario_preset_registry.py @@ -0,0 +1,206 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the scenario preset registry.""" + +import logging +from pathlib import Path + +import pytest + +from pyrit.exceptions.exception_classes import ScenarioPresetConflictError +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry +from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage + + +def _make_preset(**overrides: object) -> ScenarioPreset: + """ + Build a preset with test defaults. + + Returns: + ScenarioPreset: The constructed preset. + """ + fields: dict[str, object] = {"name": "nightly", "scenario_name": "foundry.red_team_agent"} + fields.update(overrides) + return ScenarioPreset(**fields) # type: ignore[arg-type] + + +@pytest.fixture +def registry(tmp_path: Path) -> ScenarioPresetRegistry: + """ + Build a registry backed by an isolated local directory. + + Returns: + ScenarioPresetRegistry: The registry under test. + """ + return ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(tmp_path))) + + +def test_list_presets_returns_union_sorted_by_name(registry: ScenarioPresetRegistry) -> None: + """Test that the registry exposes code-registered and file-loaded presets together.""" + registry.register_builtin(_make_preset(name="builtin_suite")) + registry.save_preset(preset=_make_preset(name="user_suite"), expected_version=None) + + names = [preset.name for preset in registry.list_presets()] + + assert names == ["builtin_suite", "user_suite"] + + +def test_register_builtin_forces_builtin_provenance(registry: ScenarioPresetRegistry) -> None: + """Test that a registered built-in cannot be marked editable by its author.""" + registered = registry.register_builtin(_make_preset(provenance=ScenarioPresetProvenance.USER)) + + assert registered.provenance is ScenarioPresetProvenance.BUILT_IN + assert registry.is_builtin("nightly") is True + + +def test_register_builtin_rejects_duplicate_name(registry: ScenarioPresetRegistry) -> None: + """Test that two built-ins cannot claim the same name.""" + registry.register_builtin(_make_preset()) + + with pytest.raises(ValueError, match="already registered"): + registry.register_builtin(_make_preset()) + + +def test_get_preset_resolves_both_provenances(registry: ScenarioPresetRegistry) -> None: + """Test name resolution across both sources.""" + registry.register_builtin(_make_preset(name="builtin_suite")) + registry.save_preset(preset=_make_preset(name="user_suite"), expected_version=None) + + assert registry.get_preset("builtin_suite") is not None + assert registry.get_preset("user_suite") is not None + assert registry.get_preset("absent") is None + + +def test_saving_over_a_builtin_name_is_rejected(registry: ScenarioPresetRegistry) -> None: + """Test that forking a built-in requires a new name.""" + registry.register_builtin(_make_preset()) + + with pytest.raises(ValueError, match="built in and cannot be saved"): + registry.save_preset(preset=_make_preset(), expected_version=None) + + +def test_deleting_a_builtin_is_rejected(registry: ScenarioPresetRegistry) -> None: + """Test that built-in presets cannot be removed.""" + registry.register_builtin(_make_preset()) + + with pytest.raises(ValueError, match="built in and cannot be deleted"): + registry.delete_preset("nightly") + + +def test_fork_under_a_new_name_succeeds(registry: ScenarioPresetRegistry) -> None: + """Test that the supported customization path works.""" + registry.register_builtin(_make_preset(name="builtin_suite", max_dataset_size=200)) + + forked = registry.save_preset( + preset=_make_preset(name="builtin_suite_fork", max_dataset_size=20), expected_version=None + ) + + assert forked.name == "builtin_suite_fork" + assert forked.max_dataset_size == 20 + assert registry.get_preset("builtin_suite") is not None + + +def test_builtin_wins_collision_and_user_preset_is_skipped_with_warning( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + """Test the upgrade case: a new built-in shadows a stored preset without failing startup.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(description="user version"), expected_version=None) + + registry = ScenarioPresetRegistry(storage=storage) + registry.register_builtin(_make_preset(description="shipped version")) + + with caplog.at_level(logging.WARNING): + registry.load_stored_presets() + + resolved = registry.get_preset("nightly") + assert resolved is not None + assert resolved.description == "shipped version" + assert resolved.is_builtin is True + assert [preset.name for preset in registry.list_presets()] == ["nightly"] + assert "a built-in preset already uses that name" in caplog.text + + +def test_load_stored_presets_keeps_non_colliding_presets(tmp_path: Path) -> None: + """Test that a collision skips only the colliding preset.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(name="nightly"), expected_version=None) + storage.save_preset(preset=_make_preset(name="weekly"), expected_version=None) + + registry = ScenarioPresetRegistry(storage=storage) + registry.register_builtin(_make_preset(name="nightly")) + registry.load_stored_presets() + + assert [preset.name for preset in registry.list_presets()] == ["nightly", "weekly"] + weekly = registry.get_preset("weekly") + assert weekly is not None + assert weekly.is_builtin is False + + +def test_load_stored_presets_replaces_prior_state(tmp_path: Path) -> None: + """Test that reloading reflects presets deleted outside the registry.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + registry = ScenarioPresetRegistry(storage=storage) + registry.save_preset(preset=_make_preset(), expected_version=None) + + storage.delete_preset("nightly") + registry.load_stored_presets() + + assert registry.get_preset("nightly") is None + + +def test_builtin_presets_are_never_written_to_storage(tmp_path: Path) -> None: + """Test that registering a built-in does not persist anything.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + registry = ScenarioPresetRegistry(storage=storage) + + registry.register_builtin(_make_preset()) + + assert storage.list_presets() == {} + assert list(tmp_path.glob("*.json")) == [] + + +def test_save_preset_refreshes_the_in_memory_copy(registry: ScenarioPresetRegistry) -> None: + """Test that a save is visible through the registry without reloading.""" + first = registry.save_preset(preset=_make_preset(description="first"), expected_version=None) + registry.save_preset(preset=_make_preset(description="second"), expected_version=first.version) + + resolved = registry.get_preset("nightly") + + assert resolved is not None + assert resolved.description == "second" + assert resolved.version == 2 + + +def test_stale_save_propagates_conflict(registry: ScenarioPresetRegistry) -> None: + """Test that the registry surfaces the storage concurrency check.""" + registry.save_preset(preset=_make_preset(), expected_version=None) + + with pytest.raises(ScenarioPresetConflictError): + registry.save_preset(preset=_make_preset(description="stale"), expected_version=99) + + +def test_delete_preset_removes_from_registry_and_storage(tmp_path: Path) -> None: + """Test that deletion clears both the cache and the backing file.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + registry = ScenarioPresetRegistry(storage=storage) + registry.save_preset(preset=_make_preset(), expected_version=None) + + registry.delete_preset("nightly") + + assert registry.get_preset("nightly") is None + assert storage.list_presets() == {} + + +def test_configure_storage_source_switches_backend(tmp_path: Path) -> None: + """Test that the storage source can be pointed at an explicit directory.""" + registry = ScenarioPresetRegistry() + presets_dir = tmp_path / "presets" + presets_dir.mkdir() + + registry.configure_storage_source(str(presets_dir)) + registry.save_preset(preset=_make_preset(), expected_version=None) + + assert (presets_dir / "nightly.json").is_file() diff --git a/tests/unit/registry/test_scenario_preset_storage.py b/tests/unit/registry/test_scenario_preset_storage.py new file mode 100644 index 0000000000..950d0a3726 --- /dev/null +++ b/tests/unit/registry/test_scenario_preset_storage.py @@ -0,0 +1,286 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for scenario preset storage.""" + +import json +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from pyrit.exceptions.exception_classes import ScenarioPresetConflictError +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage + + +def _make_preset(**overrides: object) -> ScenarioPreset: + """ + Build a preset with test defaults. + + Returns: + ScenarioPreset: The constructed preset. + """ + fields: dict[str, object] = {"name": "nightly", "scenario_name": "foundry.red_team_agent"} + fields.update(overrides) + return ScenarioPreset(**fields) # type: ignore[arg-type] + + +def test_local_storage_round_trips_a_preset(tmp_path: Path) -> None: + """Test that a saved preset loads back with its fields intact.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + preset = _make_preset( + techniques=["crescendo"], + dataset_names=["harmbench"], + max_dataset_size=25, + dataset_filters={"harm_categories": ["violence"]}, + include_baseline=True, + scenario_params={"max_turns": 3}, + description="Nightly smoke suite", + ) + + storage.save_preset(preset=preset, expected_version=None) + loaded = storage.load_preset("nightly") + + assert loaded is not None + assert loaded.techniques == ["crescendo"] + assert loaded.dataset_names == ["harmbench"] + assert loaded.max_dataset_size == 25 + assert loaded.dataset_filters == {"harm_categories": ["violence"]} + assert loaded.include_baseline is True + assert loaded.scenario_params == {"max_turns": 3} + assert loaded.description == "Nightly smoke suite" + + +def test_unset_fields_round_trip_as_none_not_false(tmp_path: Path) -> None: + """Test the tri-state regression: an unset override must not deserialize as a disabled one.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + storage.save_preset(preset=_make_preset(), expected_version=None) + loaded = storage.load_preset("nightly") + + assert loaded is not None + assert loaded.include_baseline is None + assert loaded.max_dataset_size is None + assert loaded.techniques is None + + stored = json.loads((tmp_path / "nightly.json").read_text(encoding="utf-8")) + assert "include_baseline" not in stored + assert "max_dataset_size" not in stored + + +def test_explicit_false_round_trips_as_false(tmp_path: Path) -> None: + """Test that an explicitly disabled override survives storage.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + storage.save_preset(preset=_make_preset(include_baseline=False), expected_version=None) + loaded = storage.load_preset("nightly") + + assert loaded is not None + assert loaded.include_baseline is False + + +def test_create_assigns_version_one(tmp_path: Path) -> None: + """Test that creating a preset starts its change counter at one.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + saved = storage.save_preset(preset=_make_preset(version=99), expected_version=None) + + assert saved.version == 1 + + +def test_update_increments_version(tmp_path: Path) -> None: + """Test that each accepted save advances the change counter.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + first = storage.save_preset(preset=_make_preset(), expected_version=None) + + second = storage.save_preset(preset=_make_preset(description="changed"), expected_version=first.version) + third = storage.save_preset(preset=_make_preset(description="again"), expected_version=second.version) + + assert (first.version, second.version, third.version) == (1, 2, 3) + + +def test_stale_version_save_is_rejected(tmp_path: Path) -> None: + """Test that a save based on a superseded version loses the concurrency check.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(), expected_version=None) + storage.save_preset(preset=_make_preset(description="first writer"), expected_version=1) + + with pytest.raises(ScenarioPresetConflictError) as error: + storage.save_preset(preset=_make_preset(description="second writer"), expected_version=1) + + assert error.value.expected_version == 1 + assert error.value.actual_version == 2 + assert error.value.status_code == 409 + + +def test_stale_save_does_not_overwrite_the_winner(tmp_path: Path) -> None: + """Test that a rejected save leaves stored content untouched.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(description="original"), expected_version=None) + + with pytest.raises(ScenarioPresetConflictError): + storage.save_preset(preset=_make_preset(description="clobber"), expected_version=99) + + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.description == "original" + + +def test_create_over_existing_name_is_rejected(tmp_path: Path) -> None: + """Test that creating a preset that already exists is a conflict, not an overwrite.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(), expected_version=None) + + with pytest.raises(ScenarioPresetConflictError) as error: + storage.save_preset(preset=_make_preset(description="second"), expected_version=None) + + assert error.value.expected_version is None + assert error.value.actual_version == 1 + + +def test_update_of_missing_preset_is_rejected(tmp_path: Path) -> None: + """Test that updating a preset deleted by someone else is a conflict.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + with pytest.raises(ScenarioPresetConflictError) as error: + storage.save_preset(preset=_make_preset(), expected_version=3) + + assert error.value.actual_version is None + + +def test_save_forces_user_provenance(tmp_path: Path) -> None: + """Test that stored presets are always user-owned regardless of the submitted value.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + saved = storage.save_preset(preset=_make_preset(), expected_version=None) + + assert saved.provenance is ScenarioPresetProvenance.USER + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.provenance is ScenarioPresetProvenance.USER + + +def test_builtin_preset_is_never_written(tmp_path: Path) -> None: + """Test that a built-in preset cannot be persisted to user storage.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + builtin = _make_preset(provenance=ScenarioPresetProvenance.BUILT_IN) + + with pytest.raises(ValueError, match="built in and cannot be saved"): + storage.save_preset(preset=builtin, expected_version=None) + + assert list(tmp_path.glob("*.json")) == [] + + +def test_load_missing_preset_returns_none(tmp_path: Path) -> None: + """Test that an absent preset loads as None rather than raising.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + assert storage.load_preset("absent") is None + + +def test_list_and_delete_presets(tmp_path: Path) -> None: + """Test the storage lifecycle for multiple presets.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(name="nightly"), expected_version=None) + storage.save_preset(preset=_make_preset(name="weekly"), expected_version=None) + + assert sorted(storage.list_presets()) == ["nightly", "weekly"] + storage.delete_preset("nightly") + assert sorted(storage.list_presets()) == ["weekly"] + + +def test_delete_missing_preset_is_silent(tmp_path: Path) -> None: + """Test that deleting an absent preset is not an error.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + storage.delete_preset("absent") + + +def test_malformed_preset_is_skipped_not_fatal(tmp_path: Path) -> None: + """Test that one unparseable file does not prevent the rest of the library from loading.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(name="valid"), expected_version=None) + (tmp_path / "broken.json").write_text("{not json", encoding="utf-8") + (tmp_path / "wrong_shape.json").write_text('["a list"]', encoding="utf-8") + + presets = storage.list_presets() + + assert sorted(presets) == ["valid"] + assert storage.load_preset("broken") is None + + +def test_document_name_overrides_payload_name(tmp_path: Path) -> None: + """Test that the file name is authoritative, so a load and its later save agree on the key.""" + (tmp_path / "actual_key.json").write_text( + json.dumps({"name": "different", "scenario_name": "foundry.red_team_agent", "version": 1}), + encoding="utf-8", + ) + storage = ScenarioPresetStorage(source=str(tmp_path)) + + loaded = storage.load_preset("actual_key") + + assert loaded is not None + assert loaded.name == "actual_key" + assert sorted(storage.list_presets()) == ["actual_key"] + + +def test_local_storage_returns_preset_path(tmp_path: Path) -> None: + """Test resolving the displayed path for a local preset.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + assert storage.get_preset_source("nightly") == str(tmp_path / "nightly.json") + assert storage.display_source == str(tmp_path) + + +@pytest.mark.parametrize( + "source", + [ + "https://account.blob.attacker.example/presets", + "https://user@account.blob.core.windows.net/presets", + "https://blob.core.windows.net/presets", + ], +) +def test_blob_storage_rejects_untrusted_authorities(source: str) -> None: + """Test rejecting Blob lookalikes before Azure credentials are acquired.""" + with pytest.raises(ValueError, match="local directory or Azure Blob container URI"): + ScenarioPresetStorage(source=source) + + +def test_blob_storage_round_trips_and_ignores_other_extensions() -> None: + """Test container storage operations and that only JSON documents are listed.""" + document = json.dumps({"name": "nightly", "scenario_name": "foundry.red_team_agent", "version": 4}) + client = MagicMock() + client.__enter__.return_value = client + client.list_blobs.return_value = [ + SimpleNamespace(name="nightly.json"), + SimpleNamespace(name="notes.txt"), + SimpleNamespace(name="archive/ignored.json"), + ] + client.download_blob.return_value.readall.return_value = document.encode("utf-8") + source = "https://account.blob.core.windows.net/presets?sp=rwd&sig=secret" + storage = ScenarioPresetStorage(source=source) + + with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): + presets = storage.list_presets() + saved = storage.save_preset(preset=_make_preset(description="updated"), expected_version=4) + + assert sorted(presets) == ["nightly"] + assert presets["nightly"].version == 4 + assert saved.version == 5 + assert storage.display_source == "https://account.blob.core.windows.net/presets" + assert client.upload_blob.call_args.kwargs["name"] == "nightly.json" + + +def test_blob_storage_reads_missing_document_as_none() -> None: + """Test that a missing blob loads as None instead of propagating an Azure error.""" + from azure.core.exceptions import ResourceNotFoundError + + client = MagicMock() + client.__enter__.return_value = client + client.download_blob.side_effect = ResourceNotFoundError("missing") + storage = ScenarioPresetStorage(source="https://account.blob.core.windows.net/presets?sig=secret") + + with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): + assert storage.load_preset("absent") is None From ed0fa4bd1c99cd4917e1f3e9f1e9609c3e11dfcc Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Thu, 1 Oct 2026 16:59:05 -0400 Subject: [PATCH 02/16] FIX: Harden preset storage names, versions, and registry reads Addresses review feedback on the preset storage contract. - Validate registry names in every single-document operation on FileDocumentStorage, closing a path-traversal hole that let a preset name containing separators read, overwrite, or delete files outside the configured source. The check lives at the base so no document API can omit it, which hardens CustomInitializerStorage as well. - Replace the in-model monotonic version counter with an opaque token derived from the stored bytes, paired with the preset as StoredPreset. The counter lived inside the document it guarded, so a hand-edit rewrote the very value used to detect that edit. Hashing the raw content also makes save_preset refuse to create over a malformed file instead of clobbering it. - Make ScenarioPresetRegistry read user presets through to storage instead of caching them, removing staleness across processes and the non-atomic cache rebind. load_stored_presets() is gone; get_stored_preset() replaces it. - Move ScenarioPresetConflictError out of pyrit.exceptions into the storage module as a ValueError subclass and drop its inert status_code. - Drop the misleading is_builtin guard in save_preset, which claimed to protect built-ins while acting on caller-supplied provenance. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/exceptions/__init__.py | 2 - pyrit/exceptions/exception_classes.py | 33 ---- pyrit/models/__init__.py | 2 + pyrit/models/catalog/__init__.py | 3 +- pyrit/models/catalog/scenario_preset.py | 17 +- pyrit/registry/__init__.py | 3 +- pyrit/registry/file_document_storage.py | 30 +++- pyrit/registry/scenario_preset_registry.py | 84 ++++++---- pyrit/registry/scenario_preset_storage.py | 114 ++++++++----- tests/unit/models/test_scenario_preset.py | 8 +- .../registry/test_scenario_preset_registry.py | 75 +++++++-- .../registry/test_scenario_preset_storage.py | 154 ++++++++++++------ 12 files changed, 344 insertions(+), 181 deletions(-) diff --git a/pyrit/exceptions/__init__.py b/pyrit/exceptions/__init__.py index dbae736cbe..a42582594b 100644 --- a/pyrit/exceptions/__init__.py +++ b/pyrit/exceptions/__init__.py @@ -22,7 +22,6 @@ PyritException, RateLimitException, ScenarioPartialFailureException, - ScenarioPresetConflictError, ScorerLLMResponseBlockedException, get_retry_max_num_attempts, handle_bad_request_exception, @@ -78,7 +77,6 @@ "remove_markdown_json": "pyrit.exceptions.exceptions_helpers", "RetryCollector": "pyrit.exceptions.retry_collector", "ScenarioPartialFailureException": "pyrit.exceptions.exception_classes", - "ScenarioPresetConflictError": "pyrit.exceptions.exception_classes", "ScorerLLMResponseBlockedException": "pyrit.exceptions.exception_classes", "set_execution_context": "pyrit.exceptions.exception_context", "set_retry_collector": "pyrit.exceptions.retry_collector", diff --git a/pyrit/exceptions/exception_classes.py b/pyrit/exceptions/exception_classes.py index 483fe4f70e..b15fa8272c 100644 --- a/pyrit/exceptions/exception_classes.py +++ b/pyrit/exceptions/exception_classes.py @@ -302,39 +302,6 @@ def __init__( self.__cause__ = self.incomplete_objectives[0][1] -class ScenarioPresetConflictError(PyritException): - """ - Exception raised when a scenario preset save loses an optimistic-concurrency check. - - Carries the version the caller edited against and the version currently stored so a - caller can show the user what changed underneath them rather than silently overwriting. - """ - - def __init__(self, *, name: str, expected_version: int | None, actual_version: int | None) -> None: - """ - Initialize a scenario preset conflict error. - - Args: - name (str): Name of the preset that could not be saved. - expected_version (int | None): Version the caller based its edit on, or ``None`` - when the caller intended to create a new preset. - actual_version (int | None): Version currently stored, or ``None`` when no preset - with that name exists. - """ - self.name = name - self.expected_version = expected_version - self.actual_version = actual_version - - if actual_version is None: - detail = f"expected version {expected_version} but it no longer exists" - elif expected_version is None: - detail = f"it already exists at version {actual_version}" - else: - detail = f"expected version {expected_version} but found {actual_version}" - - super().__init__(status_code=409, message=f"Scenario preset '{name}' could not be saved: {detail}.") - - class InvalidJsonException(PyritException): """Exception class for blocked content errors.""" diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index f5342a6996..0faf445791 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -59,6 +59,7 @@ ScenarioRunSizeEstimateStatus, ScenarioRunSizeFactor, ScenarioTechniqueSummary, + StoredPreset, ) from pyrit.models.conversation_stats import ConversationStats from pyrit.models.embeddings import EmbeddingData, EmbeddingResponse, EmbeddingSupport, EmbeddingUsageInformation @@ -414,6 +415,7 @@ "warn_prompt_path_deprecated": "pyrit.models.seeds", "snake_case_to_class_name": "pyrit.models.identifiers", "sort_message_pieces": "pyrit.models.messages.message_piece", + "StoredPreset": "pyrit.models.catalog", "StrategyResult": "pyrit.models.results.strategy_result", "StrategyResultT": "pyrit.models.results.strategy_result", "StructuredParameterValue": "pyrit.models.parameter", diff --git a/pyrit/models/catalog/__init__.py b/pyrit/models/catalog/__init__.py index f0739608f5..532f1ae08f 100644 --- a/pyrit/models/catalog/__init__.py +++ b/pyrit/models/catalog/__init__.py @@ -38,7 +38,7 @@ ScenarioRunSummary, ScenarioTechniqueSummary, ) - from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance + from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset from pyrit.models.catalog.target import TargetInstance _LAZY_EXPORTS: dict[str, str] = { @@ -61,6 +61,7 @@ "ScenarioRunSizeFactor": "pyrit.models.catalog.scenario", "ScenarioRunSummary": "pyrit.models.catalog.scenario", "ScenarioTechniqueSummary": "pyrit.models.catalog.scenario", + "StoredPreset": "pyrit.models.catalog.scenario_preset", "TargetInstance": "pyrit.models.catalog.target", } diff --git a/pyrit/models/catalog/scenario_preset.py b/pyrit/models/catalog/scenario_preset.py index 8ac955b470..efa978c349 100644 --- a/pyrit/models/catalog/scenario_preset.py +++ b/pyrit/models/catalog/scenario_preset.py @@ -55,7 +55,6 @@ class ScenarioPreset(BaseModel): scenario_params: dict[str, Any] | None = Field( None, description="Scenario-declared parameters such as template names and attempt counts" ) - version: int = Field(1, ge=1, description="Monotonic change counter assigned by storage, not an address") provenance: ScenarioPresetProvenance = Field( ScenarioPresetProvenance.USER, description="Whether this preset ships with PyRIT or was authored by a user" ) @@ -90,3 +89,19 @@ def _validate_dataset_filters(cls, value: dict[str, list[str]] | None) -> dict[s dict[str, list[str]] | None: Validated filters. """ return _validate_dataset_filter_mapping(value) + + +class StoredPreset(BaseModel): + """ + A preset together with the version of the document it was read from. + + The version describes the *stored document*, not the preset, so it is paired with + the preset rather than carried as a field on it. Storing the token inside the file + it guards would let a hand-edit rewrite the very value used to detect that edit. + + The token is opaque. Callers round-trip it from a read back into a write and must + not parse, compare, or order it. + """ + + preset: ScenarioPreset = Field(..., description="The stored preset") + version: str = Field(..., description="Opaque version of the document this preset was read from") diff --git a/pyrit/registry/__init__.py b/pyrit/registry/__init__.py index d6bb1afb58..5619fafac7 100644 --- a/pyrit/registry/__init__.py +++ b/pyrit/registry/__init__.py @@ -33,7 +33,7 @@ from pyrit.registry.registry import InstanceHoldingRegistry, ParamBagRegistry, Registry from pyrit.registry.registry_metadata import RegistryMetadata from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry - from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage + from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage from pyrit.registry.tag_query import TagQuery _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { @@ -53,6 +53,7 @@ "InitializerRegistry": "pyrit.registry.components", "RegistryEntry": "pyrit.registry.instance_registry", "ScenarioMetadata": "pyrit.registry.components", + "ScenarioPresetConflictError": "pyrit.registry.scenario_preset_storage", "ScenarioPresetRegistry": "pyrit.registry.scenario_preset_registry", "ScenarioPresetStorage": "pyrit.registry.scenario_preset_storage", "ScenarioRegistry": "pyrit.registry.components", diff --git a/pyrit/registry/file_document_storage.py b/pyrit/registry/file_document_storage.py index 25cade9d78..c6cd56e57d 100644 --- a/pyrit/registry/file_document_storage.py +++ b/pyrit/registry/file_document_storage.py @@ -11,6 +11,7 @@ from urllib.parse import unquote, urlparse from pyrit.common.azure_storage import has_sas_signature, is_azure_blob_uri, redact_url_credentials +from pyrit.models.identifiers.class_name_utils import validate_registry_name if TYPE_CHECKING: from collections.abc import Generator @@ -26,6 +27,11 @@ class FileDocumentStorage: fixed extension. Nested paths are ignored so a virtual directory prefix behaves the same way in both backends. + Every operation that addresses a single document validates the name first. The + name becomes a path component in both backends, so an unvalidated name would let + a caller read, overwrite, or delete a file outside the configured source. The + check lives here rather than in each subclass so no document API can omit it. + Subclasses supply the extension and a human-readable label for error messages, then expose a domain-specific API over the protected document operations. """ @@ -67,7 +73,11 @@ def _get_document_source(self, name: str) -> str: Returns: str: Local file path or Azure Blob URI for the document. + + Raises: + ValueError: If *name* is not a legal registry name. """ + validate_registry_name(name) if self._is_blob: return f"{self.display_source.rstrip('/')}/{name}{self._extension}" return str(self._local_directory() / f"{name}{self._extension}") @@ -91,7 +101,11 @@ def _read_document(self, name: str) -> str | None: Returns: str | None: Document content, or ``None`` if it does not exist. + + Raises: + ValueError: If *name* is not a legal registry name. """ + validate_registry_name(name) if self._is_blob: from azure.core.exceptions import ResourceNotFoundError @@ -105,7 +119,13 @@ def _read_document(self, name: str) -> str | None: return path.read_text(encoding="utf-8") if path.is_file() else None def _save_document(self, *, name: str, content: str) -> None: - """Persist one document, overwriting any existing content.""" + """ + Persist one document, overwriting any existing content. + + Raises: + ValueError: If *name* is not a legal registry name. + """ + validate_registry_name(name) if self._is_blob: with self._open_container_client() as client: client.upload_blob(name=self._get_blob_name(name), data=content.encode("utf-8"), overwrite=True) @@ -114,7 +134,13 @@ def _save_document(self, *, name: str, content: str) -> None: (directory / f"{name}{self._extension}").write_text(content, encoding="utf-8") def _delete_document(self, name: str) -> None: - """Delete one document if it exists.""" + """ + Delete one document if it exists. + + Raises: + ValueError: If *name* is not a legal registry name. + """ + validate_registry_name(name) if self._is_blob: from azure.core.exceptions import ResourceNotFoundError diff --git a/pyrit/registry/scenario_preset_registry.py b/pyrit/registry/scenario_preset_registry.py index 1ff2c8c82e..174d345760 100644 --- a/pyrit/registry/scenario_preset_registry.py +++ b/pyrit/registry/scenario_preset_registry.py @@ -1,14 +1,14 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -"""In-memory registry unioning built-in and user-authored scenario presets.""" +"""Registry unioning built-in and user-authored scenario presets.""" from __future__ import annotations import logging from typing import TYPE_CHECKING -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage if TYPE_CHECKING: @@ -22,8 +22,11 @@ class ScenarioPresetRegistry: The union of built-in and user-authored scenario presets. Presets arrive from two places. Built-in presets are registered from initializer code - at startup and are read-only. User presets are loaded from storage. The registry exists - because only it sees both, so only it can resolve a name or detect a dangling reference. + at startup and are read-only, so they are held in memory. User presets are read from + storage on each call rather than cached, because the same directory or blob container + is routinely shared between a notebook, the API, and a second process; a cache would + serve edits those callers can no longer see. The registry exists because only it sees + both sources, so only it can resolve a name across them. Built-in presets win on a name collision. A user preset that collides is skipped with a warning rather than failing startup, because the collision arises when PyRIT ships a new @@ -41,7 +44,6 @@ def __init__(self, *, storage: ScenarioPresetStorage | None = None) -> None: """ self._storage = storage self._builtin_presets: dict[str, ScenarioPreset] = {} - self._user_presets: dict[str, ScenarioPreset] = {} def configure_storage_source(self, source: str | None) -> None: """Configure the local directory or Azure Blob source for user presets.""" @@ -68,26 +70,37 @@ def register_builtin(self, preset: ScenarioPreset) -> ScenarioPreset: self._builtin_presets[registered.name] = registered return registered - def load_stored_presets(self) -> None: - """Load user presets from storage, skipping any that collide with a built-in preset.""" - self._user_presets = {} - for name, preset in self._get_storage().list_presets().items(): - if name in self._builtin_presets: - logger.warning( - f"Skipping stored scenario preset '{name}': a built-in preset already uses that name. " - "Rename the stored preset to keep using it." - ) - continue - self._user_presets[name] = preset - def get_preset(self, name: str) -> ScenarioPreset | None: """ Resolve one preset by name. Returns: ScenarioPreset | None: The preset, or ``None`` if no preset uses that name. + + Raises: + ValueError: If *name* is not a legal preset name. + """ + builtin = self._builtin_presets.get(name) + if builtin is not None: + return builtin + + stored = self._get_storage().load_preset(name) + return stored.preset if stored is not None else None + + def get_stored_preset(self, name: str) -> StoredPreset | None: """ - return self._builtin_presets.get(name) or self._user_presets.get(name) + Read one user preset together with the version needed to save an edit to it. + + Returns: + StoredPreset | None: The stored preset, or ``None`` if *name* is built in or + no stored preset uses that name. + + Raises: + ValueError: If *name* is not a legal preset name. + """ + if name in self._builtin_presets: + return None + return self._get_storage().load_preset(name) def list_presets(self) -> list[ScenarioPreset]: """ @@ -96,46 +109,51 @@ def list_presets(self) -> list[ScenarioPreset]: Returns: list[ScenarioPreset]: Built-in and user presets, sorted by name. """ - merged = {**self._user_presets, **self._builtin_presets} + merged: dict[str, ScenarioPreset] = {} + for name, stored in self._get_storage().list_presets().items(): + if name in self._builtin_presets: + logger.warning( + f"Skipping stored scenario preset '{name}': a built-in preset already uses that name. " + "Rename the stored preset to keep using it." + ) + continue + merged[name] = stored.preset + + merged.update(self._builtin_presets) return [merged[name] for name in sorted(merged)] def is_builtin(self, name: str) -> bool: """Return whether *name* belongs to a built-in preset.""" return name in self._builtin_presets - def save_preset(self, *, preset: ScenarioPreset, expected_version: int | None) -> ScenarioPreset: + def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) -> StoredPreset: """ - Persist a user preset and refresh the in-memory copy. + Persist a user preset. Args: preset (ScenarioPreset): The preset to persist. - expected_version (int | None): ``None`` to create a preset that must not already - exist, or the version the edit was based on. + expected_version (str | None): ``None`` to create a preset that must not already + exist, or the version returned when the edited preset was read. Returns: - ScenarioPreset: The persisted preset, carrying its newly assigned version. + StoredPreset: The persisted preset and its new version. Raises: ScenarioPresetConflictError: If the stored version does not match *expected_version*. - ValueError: If *name* belongs to a built-in preset. + ValueError: If the name belongs to a built-in preset or is not a legal preset name. """ self._reject_builtin(preset.name, action="saved") - - saved = self._get_storage().save_preset(preset=preset, expected_version=expected_version) - self._user_presets[saved.name] = saved - return saved + return self._get_storage().save_preset(preset=preset, expected_version=expected_version) def delete_preset(self, name: str) -> None: """ - Delete a user preset from storage and from the registry. + Delete a user preset from storage. Raises: - ValueError: If *name* belongs to a built-in preset. + ValueError: If *name* belongs to a built-in preset or is not a legal preset name. """ self._reject_builtin(name, action="deleted") - self._get_storage().delete_preset(name) - self._user_presets.pop(name, None) def _reject_builtin(self, name: str, *, action: str) -> None: """ diff --git a/pyrit/registry/scenario_preset_storage.py b/pyrit/registry/scenario_preset_storage.py index 59984945c5..bce730045f 100644 --- a/pyrit/registry/scenario_preset_storage.py +++ b/pyrit/registry/scenario_preset_storage.py @@ -5,16 +5,43 @@ from __future__ import annotations +import hashlib import json import logging -from pyrit.exceptions.exception_classes import ScenarioPresetConflictError -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset from pyrit.registry.file_document_storage import FileDocumentStorage logger = logging.getLogger(__name__) +class ScenarioPresetConflictError(ValueError): + """A stored preset changed after the caller read it.""" + + def __init__(self, *, name: str, expected_version: str | None, actual_version: str | None) -> None: + """Initialize the error with the versions that failed to match.""" + self.name = name + self.expected_version = expected_version + self.actual_version = actual_version + if actual_version is None: + detail = "it no longer exists" + elif expected_version is None: + detail = "it already exists" + else: + detail = "it was changed by someone else" + super().__init__(f"Scenario preset '{name}' could not be saved because {detail}. Reload it and reapply.") + + +def _document_version(content: str) -> str: + """ + Create an opaque version token from stored document content. + + Returns: + str: The document-state version token. + """ + return hashlib.sha256(content.encode()).hexdigest() + + class ScenarioPresetStorage(FileDocumentStorage): """ Read and write user-authored scenario presets as JSON documents. @@ -22,14 +49,15 @@ class ScenarioPresetStorage(FileDocumentStorage): Only user presets are stored. Built-in presets come from initializer code and are never written here, so a file on disk always means a user authored it. - Writes are guarded by an optimistic-concurrency check on ``version``. The stored - version is authoritative: a caller supplies the version it edited against, and the - save is refused if storage has moved on. + Writes are guarded by an optimistic-concurrency check. A caller supplies the version + it read, and the save is refused unless storage still holds that version. The token + is a hash of the stored bytes, so it also catches a file edited by hand or by another + process rather than only writes made through this class. The check is read-then-write rather than a true compare-and-swap, so two saves racing within the same instant can both observe the same version and the later write wins. - It is aimed at the realistic case — a person editing a copy that went stale minutes - ago — not at concurrent writers. Closing that gap needs backend-specific conditional + It is aimed at the realistic case - a person editing a copy that went stale minutes + ago - not at concurrent writers. Closing that gap needs backend-specific conditional writes (blob ETags have no local-filesystem equivalent) and is deliberately deferred. """ @@ -48,10 +76,13 @@ def get_preset_source(self, name: str) -> str: Returns: str: Local file path or Azure Blob URI for the preset. + + Raises: + ValueError: If *name* is not a legal preset name. """ return self._get_document_source(name) - def list_presets(self) -> dict[str, ScenarioPreset]: + def list_presets(self) -> dict[str, StoredPreset]: """ Read every stored preset, skipping any that cannot be parsed. @@ -59,71 +90,76 @@ def list_presets(self) -> dict[str, ScenarioPreset]: loading, so failures are logged and that preset is omitted. Returns: - dict[str, ScenarioPreset]: Presets keyed by name. + dict[str, StoredPreset]: Stored presets keyed by name. """ - presets: dict[str, ScenarioPreset] = {} + presets: dict[str, StoredPreset] = {} for name, content in self._list_documents().items(): preset = self._parse_preset(name=name, content=content) if preset is not None: - presets[preset.name] = preset + presets[preset.name] = StoredPreset(preset=preset, version=_document_version(content)) return presets - def load_preset(self, name: str) -> ScenarioPreset | None: + def load_preset(self, name: str) -> StoredPreset | None: """ Read one stored preset. Returns: - ScenarioPreset | None: The preset, or ``None`` if it is absent or malformed. + StoredPreset | None: The preset and its version, or ``None`` if it is absent + or malformed. + + Raises: + ValueError: If *name* is not a legal preset name. """ content = self._read_document(name) if content is None: return None - return self._parse_preset(name=name, content=content) + preset = self._parse_preset(name=name, content=content) + if preset is None: + return None + return StoredPreset(preset=preset, version=_document_version(content)) - def save_preset(self, *, preset: ScenarioPreset, expected_version: int | None) -> ScenarioPreset: + def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) -> StoredPreset: """ - Persist one preset, assigning its version. + Persist one preset. - The caller states its intent through *expected_version* rather than through the - version on *preset*, so creating and updating are never ambiguous and a + The caller states its intent through *expected_version* rather than through + anything on *preset*, so creating and updating are never ambiguous and a client-supplied version can never be trusted into storage. + A document that exists but cannot be parsed still has a version, so creating + over a malformed file conflicts rather than silently discarding it. + Args: preset (ScenarioPreset): The preset to persist. - expected_version (int | None): ``None`` to create a preset that must not already - exist, or the version the edit was based on. + expected_version (str | None): ``None`` to create a preset that must not already + exist, or the version returned when the edited preset was read. Returns: - ScenarioPreset: The persisted preset, carrying its newly assigned version. + StoredPreset: The persisted preset and its new version. Raises: ScenarioPresetConflictError: If the stored version does not match *expected_version*. - ValueError: If *preset* is built-in, which is never stored. + ValueError: If the preset name is not a legal preset name. """ - if preset.is_builtin: - raise ValueError( - f"Scenario preset '{preset.name}' is built in and cannot be saved. " - "Fork it under a new name to customize it." - ) - - existing = self.load_preset(preset.name) - actual_version = existing.version if existing is not None else None + existing_content = self._read_document(preset.name) + actual_version = None if existing_content is None else _document_version(existing_content) if actual_version != expected_version: raise ScenarioPresetConflictError( name=preset.name, expected_version=expected_version, actual_version=actual_version ) - saved = preset.model_copy( - update={ - "version": 1 if existing is None else existing.version + 1, - "provenance": ScenarioPresetProvenance.USER, - } - ) - self._save_document(name=saved.name, content=self._serialize_preset(saved)) - return saved + saved = preset.model_copy(update={"provenance": ScenarioPresetProvenance.USER}) + content = self._serialize_preset(saved) + self._save_document(name=saved.name, content=content) + return StoredPreset(preset=saved, version=_document_version(content)) def delete_preset(self, name: str) -> None: - """Delete one stored preset if it exists.""" + """ + Delete one stored preset if it exists. + + Raises: + ValueError: If *name* is not a legal preset name. + """ self._delete_document(name) @staticmethod diff --git a/tests/unit/models/test_scenario_preset.py b/tests/unit/models/test_scenario_preset.py index 480ce460c0..5e8513ab4f 100644 --- a/tests/unit/models/test_scenario_preset.py +++ b/tests/unit/models/test_scenario_preset.py @@ -21,15 +21,19 @@ def test_init_defaults_every_optional_field_to_none() -> None: assert preset.description is None -def test_init_defaults_to_user_provenance_at_version_one() -> None: +def test_init_defaults_to_user_provenance() -> None: """Test the default identity of a freshly constructed preset.""" preset = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent") - assert preset.version == 1 assert preset.provenance is ScenarioPresetProvenance.USER assert preset.is_builtin is False +def test_preset_carries_no_storage_version() -> None: + """Test that storage metadata stays out of the model it describes.""" + assert "version" not in ScenarioPreset.model_fields + + def test_is_builtin_reflects_provenance() -> None: """Test that built-in provenance marks a preset read-only.""" preset = ScenarioPreset( diff --git a/tests/unit/registry/test_scenario_preset_registry.py b/tests/unit/registry/test_scenario_preset_registry.py index 6a8719d2fc..bf96bd645c 100644 --- a/tests/unit/registry/test_scenario_preset_registry.py +++ b/tests/unit/registry/test_scenario_preset_registry.py @@ -8,10 +8,9 @@ import pytest -from pyrit.exceptions.exception_classes import ScenarioPresetConflictError from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry -from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage +from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage def _make_preset(**overrides: object) -> ScenarioPreset: @@ -73,6 +72,30 @@ def test_get_preset_resolves_both_provenances(registry: ScenarioPresetRegistry) assert registry.get_preset("absent") is None +def test_get_preset_rejects_illegal_names(registry: ScenarioPresetRegistry) -> None: + """Test that a traversal attempt fails loudly rather than resolving to nothing.""" + with pytest.raises(ValueError, match="Invalid registry name"): + registry.get_preset("../victim") + + +def test_get_stored_preset_returns_the_version_needed_to_edit(registry: ScenarioPresetRegistry) -> None: + """Test that an editor can read a preset and the token required to save it back.""" + saved = registry.save_preset(preset=_make_preset(), expected_version=None) + + stored = registry.get_stored_preset("nightly") + + assert stored is not None + assert stored.version == saved.version + registry.save_preset(preset=_make_preset(description="edited"), expected_version=stored.version) + + +def test_get_stored_preset_returns_none_for_a_builtin(registry: ScenarioPresetRegistry) -> None: + """Test that a built-in has no stored document and therefore no editable version.""" + registry.register_builtin(_make_preset()) + + assert registry.get_stored_preset("nightly") is None + + def test_saving_over_a_builtin_name_is_rejected(registry: ScenarioPresetRegistry) -> None: """Test that forking a built-in requires a new name.""" registry.register_builtin(_make_preset()) @@ -97,8 +120,8 @@ def test_fork_under_a_new_name_succeeds(registry: ScenarioPresetRegistry) -> Non preset=_make_preset(name="builtin_suite_fork", max_dataset_size=20), expected_version=None ) - assert forked.name == "builtin_suite_fork" - assert forked.max_dataset_size == 20 + assert forked.preset.name == "builtin_suite_fork" + assert forked.preset.max_dataset_size == 20 assert registry.get_preset("builtin_suite") is not None @@ -113,17 +136,17 @@ def test_builtin_wins_collision_and_user_preset_is_skipped_with_warning( registry.register_builtin(_make_preset(description="shipped version")) with caplog.at_level(logging.WARNING): - registry.load_stored_presets() + listed = registry.list_presets() resolved = registry.get_preset("nightly") assert resolved is not None assert resolved.description == "shipped version" assert resolved.is_builtin is True - assert [preset.name for preset in registry.list_presets()] == ["nightly"] + assert [preset.name for preset in listed] == ["nightly"] assert "a built-in preset already uses that name" in caplog.text -def test_load_stored_presets_keeps_non_colliding_presets(tmp_path: Path) -> None: +def test_collision_skips_only_the_colliding_preset(tmp_path: Path) -> None: """Test that a collision skips only the colliding preset.""" storage = ScenarioPresetStorage(source=str(tmp_path)) storage.save_preset(preset=_make_preset(name="nightly"), expected_version=None) @@ -131,7 +154,6 @@ def test_load_stored_presets_keeps_non_colliding_presets(tmp_path: Path) -> None registry = ScenarioPresetRegistry(storage=storage) registry.register_builtin(_make_preset(name="nightly")) - registry.load_stored_presets() assert [preset.name for preset in registry.list_presets()] == ["nightly", "weekly"] weekly = registry.get_preset("weekly") @@ -139,14 +161,24 @@ def test_load_stored_presets_keeps_non_colliding_presets(tmp_path: Path) -> None assert weekly.is_builtin is False -def test_load_stored_presets_replaces_prior_state(tmp_path: Path) -> None: - """Test that reloading reflects presets deleted outside the registry.""" +def test_presets_written_by_another_process_are_visible(tmp_path: Path) -> None: + """Test that user presets are read through, so a shared source is not served from a stale cache.""" + registry = ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(tmp_path))) + assert registry.get_preset("nightly") is None + + ScenarioPresetStorage(source=str(tmp_path)).save_preset(preset=_make_preset(), expected_version=None) + + assert registry.get_preset("nightly") is not None + assert [preset.name for preset in registry.list_presets()] == ["nightly"] + + +def test_presets_deleted_by_another_process_disappear(tmp_path: Path) -> None: + """Test that a preset removed outside the registry stops resolving.""" storage = ScenarioPresetStorage(source=str(tmp_path)) registry = ScenarioPresetRegistry(storage=storage) registry.save_preset(preset=_make_preset(), expected_version=None) storage.delete_preset("nightly") - registry.load_stored_presets() assert registry.get_preset("nightly") is None @@ -162,7 +194,7 @@ def test_builtin_presets_are_never_written_to_storage(tmp_path: Path) -> None: assert list(tmp_path.glob("*.json")) == [] -def test_save_preset_refreshes_the_in_memory_copy(registry: ScenarioPresetRegistry) -> None: +def test_save_preset_is_immediately_visible(registry: ScenarioPresetRegistry) -> None: """Test that a save is visible through the registry without reloading.""" first = registry.save_preset(preset=_make_preset(description="first"), expected_version=None) registry.save_preset(preset=_make_preset(description="second"), expected_version=first.version) @@ -171,7 +203,6 @@ def test_save_preset_refreshes_the_in_memory_copy(registry: ScenarioPresetRegist assert resolved is not None assert resolved.description == "second" - assert resolved.version == 2 def test_stale_save_propagates_conflict(registry: ScenarioPresetRegistry) -> None: @@ -179,11 +210,11 @@ def test_stale_save_propagates_conflict(registry: ScenarioPresetRegistry) -> Non registry.save_preset(preset=_make_preset(), expected_version=None) with pytest.raises(ScenarioPresetConflictError): - registry.save_preset(preset=_make_preset(description="stale"), expected_version=99) + registry.save_preset(preset=_make_preset(description="stale"), expected_version="stale_version") def test_delete_preset_removes_from_registry_and_storage(tmp_path: Path) -> None: - """Test that deletion clears both the cache and the backing file.""" + """Test that deletion clears the backing file and stops resolution.""" storage = ScenarioPresetStorage(source=str(tmp_path)) registry = ScenarioPresetRegistry(storage=storage) registry.save_preset(preset=_make_preset(), expected_version=None) @@ -204,3 +235,17 @@ def test_configure_storage_source_switches_backend(tmp_path: Path) -> None: registry.save_preset(preset=_make_preset(), expected_version=None) assert (presets_dir / "nightly.json").is_file() + + +def test_configure_storage_source_stops_serving_the_previous_source(tmp_path: Path) -> None: + """Test that repointing the source does not leave presets from the old one resolvable.""" + first_dir = tmp_path / "first" + second_dir = tmp_path / "second" + first_dir.mkdir() + second_dir.mkdir() + registry = ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(first_dir))) + registry.save_preset(preset=_make_preset(), expected_version=None) + + registry.configure_storage_source(str(second_dir)) + + assert registry.get_preset("nightly") is None diff --git a/tests/unit/registry/test_scenario_preset_storage.py b/tests/unit/registry/test_scenario_preset_storage.py index 950d0a3726..9b8a061b94 100644 --- a/tests/unit/registry/test_scenario_preset_storage.py +++ b/tests/unit/registry/test_scenario_preset_storage.py @@ -10,9 +10,8 @@ import pytest -from pyrit.exceptions.exception_classes import ScenarioPresetConflictError from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance -from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage +from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage def _make_preset(**overrides: object) -> ScenarioPreset: @@ -44,13 +43,13 @@ def test_local_storage_round_trips_a_preset(tmp_path: Path) -> None: loaded = storage.load_preset("nightly") assert loaded is not None - assert loaded.techniques == ["crescendo"] - assert loaded.dataset_names == ["harmbench"] - assert loaded.max_dataset_size == 25 - assert loaded.dataset_filters == {"harm_categories": ["violence"]} - assert loaded.include_baseline is True - assert loaded.scenario_params == {"max_turns": 3} - assert loaded.description == "Nightly smoke suite" + assert loaded.preset.techniques == ["crescendo"] + assert loaded.preset.dataset_names == ["harmbench"] + assert loaded.preset.max_dataset_size == 25 + assert loaded.preset.dataset_filters == {"harm_categories": ["violence"]} + assert loaded.preset.include_baseline is True + assert loaded.preset.scenario_params == {"max_turns": 3} + assert loaded.preset.description == "Nightly smoke suite" def test_unset_fields_round_trip_as_none_not_false(tmp_path: Path) -> None: @@ -61,9 +60,9 @@ def test_unset_fields_round_trip_as_none_not_false(tmp_path: Path) -> None: loaded = storage.load_preset("nightly") assert loaded is not None - assert loaded.include_baseline is None - assert loaded.max_dataset_size is None - assert loaded.techniques is None + assert loaded.preset.include_baseline is None + assert loaded.preset.max_dataset_size is None + assert loaded.preset.techniques is None stored = json.loads((tmp_path / "nightly.json").read_text(encoding="utf-8")) assert "include_baseline" not in stored @@ -78,41 +77,70 @@ def test_explicit_false_round_trips_as_false(tmp_path: Path) -> None: loaded = storage.load_preset("nightly") assert loaded is not None - assert loaded.include_baseline is False + assert loaded.preset.include_baseline is False -def test_create_assigns_version_one(tmp_path: Path) -> None: - """Test that creating a preset starts its change counter at one.""" +def test_version_is_not_stored_inside_the_document(tmp_path: Path) -> None: + """Test that the conflict token lives beside the document, never inside the file it guards.""" storage = ScenarioPresetStorage(source=str(tmp_path)) - saved = storage.save_preset(preset=_make_preset(version=99), expected_version=None) + saved = storage.save_preset(preset=_make_preset(), expected_version=None) - assert saved.version == 1 + assert saved.version + assert "version" not in json.loads((tmp_path / "nightly.json").read_text(encoding="utf-8")) -def test_update_increments_version(tmp_path: Path) -> None: - """Test that each accepted save advances the change counter.""" +def test_each_save_produces_a_new_version(tmp_path: Path) -> None: + """Test that an accepted save supersedes the token the caller passed in.""" storage = ScenarioPresetStorage(source=str(tmp_path)) first = storage.save_preset(preset=_make_preset(), expected_version=None) second = storage.save_preset(preset=_make_preset(description="changed"), expected_version=first.version) - third = storage.save_preset(preset=_make_preset(description="again"), expected_version=second.version) - assert (first.version, second.version, third.version) == (1, 2, 3) + assert second.version != first.version + reloaded = storage.load_preset("nightly") + assert reloaded is not None + assert reloaded.version == second.version + + +def test_identical_content_keeps_the_same_version(tmp_path: Path) -> None: + """Test that the token describes stored content, so a no-op save does not invalidate other readers.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + first = storage.save_preset(preset=_make_preset(), expected_version=None) + + second = storage.save_preset(preset=_make_preset(), expected_version=first.version) + + assert second.version == first.version def test_stale_version_save_is_rejected(tmp_path: Path) -> None: """Test that a save based on a superseded version loses the concurrency check.""" storage = ScenarioPresetStorage(source=str(tmp_path)) - storage.save_preset(preset=_make_preset(), expected_version=None) - storage.save_preset(preset=_make_preset(description="first writer"), expected_version=1) + created = storage.save_preset(preset=_make_preset(), expected_version=None) + winner = storage.save_preset(preset=_make_preset(description="first writer"), expected_version=created.version) with pytest.raises(ScenarioPresetConflictError) as error: - storage.save_preset(preset=_make_preset(description="second writer"), expected_version=1) + storage.save_preset(preset=_make_preset(description="second writer"), expected_version=created.version) + + assert error.value.expected_version == created.version + assert error.value.actual_version == winner.version - assert error.value.expected_version == 1 - assert error.value.actual_version == 2 - assert error.value.status_code == 409 + +def test_out_of_band_edit_invalidates_the_version(tmp_path: Path) -> None: + """Test that a file edited outside this class is detected, not silently overwritten.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + created = storage.save_preset(preset=_make_preset(description="original"), expected_version=None) + (tmp_path / "nightly.json").write_text( + json.dumps({"scenario_name": "foundry.red_team_agent", "description": "edited by hand"}), + encoding="utf-8", + ) + + with pytest.raises(ScenarioPresetConflictError): + storage.save_preset(preset=_make_preset(description="stale client"), expected_version=created.version) + + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.preset.description == "edited by hand" def test_stale_save_does_not_overwrite_the_winner(tmp_path: Path) -> None: @@ -121,23 +149,34 @@ def test_stale_save_does_not_overwrite_the_winner(tmp_path: Path) -> None: storage.save_preset(preset=_make_preset(description="original"), expected_version=None) with pytest.raises(ScenarioPresetConflictError): - storage.save_preset(preset=_make_preset(description="clobber"), expected_version=99) + storage.save_preset(preset=_make_preset(description="clobber"), expected_version="not_the_stored_version") loaded = storage.load_preset("nightly") assert loaded is not None - assert loaded.description == "original" + assert loaded.preset.description == "original" def test_create_over_existing_name_is_rejected(tmp_path: Path) -> None: """Test that creating a preset that already exists is a conflict, not an overwrite.""" storage = ScenarioPresetStorage(source=str(tmp_path)) - storage.save_preset(preset=_make_preset(), expected_version=None) + created = storage.save_preset(preset=_make_preset(), expected_version=None) with pytest.raises(ScenarioPresetConflictError) as error: storage.save_preset(preset=_make_preset(description="second"), expected_version=None) assert error.value.expected_version is None - assert error.value.actual_version == 1 + assert error.value.actual_version == created.version + + +def test_create_over_malformed_file_is_rejected(tmp_path: Path) -> None: + """Test that an unreadable file still blocks a create, so hand-written content is not discarded.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + (tmp_path / "nightly.json").write_text("{not json", encoding="utf-8") + + with pytest.raises(ScenarioPresetConflictError): + storage.save_preset(preset=_make_preset(), expected_version=None) + + assert (tmp_path / "nightly.json").read_text(encoding="utf-8") == "{not json" def test_update_of_missing_preset_is_rejected(tmp_path: Path) -> None: @@ -145,7 +184,7 @@ def test_update_of_missing_preset_is_rejected(tmp_path: Path) -> None: storage = ScenarioPresetStorage(source=str(tmp_path)) with pytest.raises(ScenarioPresetConflictError) as error: - storage.save_preset(preset=_make_preset(), expected_version=3) + storage.save_preset(preset=_make_preset(), expected_version="some_version") assert error.value.actual_version is None @@ -154,23 +193,14 @@ def test_save_forces_user_provenance(tmp_path: Path) -> None: """Test that stored presets are always user-owned regardless of the submitted value.""" storage = ScenarioPresetStorage(source=str(tmp_path)) - saved = storage.save_preset(preset=_make_preset(), expected_version=None) + saved = storage.save_preset( + preset=_make_preset(provenance=ScenarioPresetProvenance.BUILT_IN), expected_version=None + ) - assert saved.provenance is ScenarioPresetProvenance.USER + assert saved.preset.provenance is ScenarioPresetProvenance.USER loaded = storage.load_preset("nightly") assert loaded is not None - assert loaded.provenance is ScenarioPresetProvenance.USER - - -def test_builtin_preset_is_never_written(tmp_path: Path) -> None: - """Test that a built-in preset cannot be persisted to user storage.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - builtin = _make_preset(provenance=ScenarioPresetProvenance.BUILT_IN) - - with pytest.raises(ValueError, match="built in and cannot be saved"): - storage.save_preset(preset=builtin, expected_version=None) - - assert list(tmp_path.glob("*.json")) == [] + assert loaded.preset.provenance is ScenarioPresetProvenance.USER def test_load_missing_preset_returns_none(tmp_path: Path) -> None: @@ -198,6 +228,25 @@ def test_delete_missing_preset_is_silent(tmp_path: Path) -> None: storage.delete_preset("absent") +@pytest.mark.parametrize("name", ["../victim", "..\\victim", "nested/victim", "Nightly", "night-ly", ""]) +def test_document_operations_reject_illegal_names(tmp_path: Path, name: str) -> None: + """Test that a name which would escape the configured source is refused on every path.""" + source = tmp_path / "presets" + source.mkdir() + victim = tmp_path / "victim.json" + victim.write_text("do not touch", encoding="utf-8") + storage = ScenarioPresetStorage(source=str(source)) + + with pytest.raises(ValueError, match="Invalid registry name"): + storage.load_preset(name) + with pytest.raises(ValueError, match="Invalid registry name"): + storage.delete_preset(name) + with pytest.raises(ValueError, match="Invalid registry name"): + storage.get_preset_source(name) + + assert victim.read_text(encoding="utf-8") == "do not touch" + + def test_malformed_preset_is_skipped_not_fatal(tmp_path: Path) -> None: """Test that one unparseable file does not prevent the rest of the library from loading.""" storage = ScenarioPresetStorage(source=str(tmp_path)) @@ -214,7 +263,7 @@ def test_malformed_preset_is_skipped_not_fatal(tmp_path: Path) -> None: def test_document_name_overrides_payload_name(tmp_path: Path) -> None: """Test that the file name is authoritative, so a load and its later save agree on the key.""" (tmp_path / "actual_key.json").write_text( - json.dumps({"name": "different", "scenario_name": "foundry.red_team_agent", "version": 1}), + json.dumps({"name": "different", "scenario_name": "foundry.red_team_agent"}), encoding="utf-8", ) storage = ScenarioPresetStorage(source=str(tmp_path)) @@ -222,7 +271,7 @@ def test_document_name_overrides_payload_name(tmp_path: Path) -> None: loaded = storage.load_preset("actual_key") assert loaded is not None - assert loaded.name == "actual_key" + assert loaded.preset.name == "actual_key" assert sorted(storage.list_presets()) == ["actual_key"] @@ -250,7 +299,7 @@ def test_blob_storage_rejects_untrusted_authorities(source: str) -> None: def test_blob_storage_round_trips_and_ignores_other_extensions() -> None: """Test container storage operations and that only JSON documents are listed.""" - document = json.dumps({"name": "nightly", "scenario_name": "foundry.red_team_agent", "version": 4}) + document = json.dumps({"name": "nightly", "scenario_name": "foundry.red_team_agent"}) client = MagicMock() client.__enter__.return_value = client client.list_blobs.return_value = [ @@ -264,11 +313,12 @@ def test_blob_storage_round_trips_and_ignores_other_extensions() -> None: with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): presets = storage.list_presets() - saved = storage.save_preset(preset=_make_preset(description="updated"), expected_version=4) + saved = storage.save_preset( + preset=_make_preset(description="updated"), expected_version=presets["nightly"].version + ) assert sorted(presets) == ["nightly"] - assert presets["nightly"].version == 4 - assert saved.version == 5 + assert saved.version != presets["nightly"].version assert storage.display_source == "https://account.blob.core.windows.net/presets" assert client.upload_blob.call_args.kwargs["name"] == "nightly.json" From f15fc7a1b6a4bfef64c33f8dbc97663fe0fa9503 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:47:42 -0400 Subject: [PATCH 03/16] REFACTOR: Defer built-in scenario presets and drop the preset registry Scenario defaults already cover most preset fields, so a built-in preset restating them duplicates shipped data. The only case built-ins uniquely serve is multiplicity, which no shipped preset needs yet, so the machinery is deferred to the PR that introduces the first real built-in preset. - Remove ScenarioPresetProvenance, provenance, and is_builtin. - Delete ScenarioPresetRegistry, which without built-ins was a pass-through to storage, and move its default-directory logic into ScenarioPresetStorage. - Re-adding provenance later is not a migration: older documents lacking the key take the default. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/models/__init__.py | 2 - pyrit/models/catalog/__init__.py | 3 +- pyrit/models/catalog/scenario_preset.py | 16 -- pyrit/registry/__init__.py | 2 - pyrit/registry/scenario_preset_registry.py | 195 -------------- pyrit/registry/scenario_preset_storage.py | 51 +++- tests/unit/models/test_scenario_preset.py | 21 +- .../registry/test_scenario_preset_registry.py | 251 ------------------ .../registry/test_scenario_preset_storage.py | 38 ++- 9 files changed, 68 insertions(+), 511 deletions(-) delete mode 100644 pyrit/registry/scenario_preset_registry.py delete mode 100644 tests/unit/registry/test_scenario_preset_registry.py diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 0faf445791..9b5fa0edcf 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -50,7 +50,6 @@ ScenarioDatasetSummary, ScenarioDefaultRunSizeEstimate, ScenarioPreset, - ScenarioPresetProvenance, ScenarioRunListItem, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, @@ -360,7 +359,6 @@ "ScenarioDatasetSummary": "pyrit.models.catalog", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog", "ScenarioPreset": "pyrit.models.catalog", - "ScenarioPresetProvenance": "pyrit.models.catalog", "ScenarioRunListItem": "pyrit.models.catalog", "ScenarioRunSizeComponent": "pyrit.models.catalog", "ScenarioRunSizeEstimate": "pyrit.models.catalog", diff --git a/pyrit/models/catalog/__init__.py b/pyrit/models/catalog/__init__.py index 532f1ae08f..d2638257bd 100644 --- a/pyrit/models/catalog/__init__.py +++ b/pyrit/models/catalog/__init__.py @@ -38,7 +38,7 @@ ScenarioRunSummary, ScenarioTechniqueSummary, ) - from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset + from pyrit.models.catalog.scenario_preset import ScenarioPreset, StoredPreset from pyrit.models.catalog.target import TargetInstance _LAZY_EXPORTS: dict[str, str] = { @@ -51,7 +51,6 @@ "ScenarioDatasetSummary": "pyrit.models.catalog.scenario", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog.scenario", "ScenarioPreset": "pyrit.models.catalog.scenario_preset", - "ScenarioPresetProvenance": "pyrit.models.catalog.scenario_preset", "ScenarioRunListItem": "pyrit.models.catalog.scenario", "ScenarioRunSizeComponent": "pyrit.models.catalog.scenario", "ScenarioRunSizeEstimate": "pyrit.models.catalog.scenario", diff --git a/pyrit/models/catalog/scenario_preset.py b/pyrit/models/catalog/scenario_preset.py index efa978c349..4583e5a94e 100644 --- a/pyrit/models/catalog/scenario_preset.py +++ b/pyrit/models/catalog/scenario_preset.py @@ -16,7 +16,6 @@ the moment a preset was saved and stop it tracking upstream changes. """ -from enum import Enum from typing import Any from pydantic import BaseModel, Field, field_validator @@ -25,13 +24,6 @@ from pyrit.models.identifiers.class_name_utils import validate_registry_name -class ScenarioPresetProvenance(str, Enum): - """Where a preset came from, which determines whether it can be edited.""" - - BUILT_IN = "built_in" - USER = "user" - - class ScenarioPreset(BaseModel): """ A named, reusable, target-agnostic scenario configuration. @@ -55,14 +47,6 @@ class ScenarioPreset(BaseModel): scenario_params: dict[str, Any] | None = Field( None, description="Scenario-declared parameters such as template names and attempt counts" ) - provenance: ScenarioPresetProvenance = Field( - ScenarioPresetProvenance.USER, description="Whether this preset ships with PyRIT or was authored by a user" - ) - - @property - def is_builtin(self) -> bool: - """Whether this preset ships with PyRIT and is therefore read-only.""" - return self.provenance is ScenarioPresetProvenance.BUILT_IN @field_validator("name") @classmethod diff --git a/pyrit/registry/__init__.py b/pyrit/registry/__init__.py index 5619fafac7..5dc9cad219 100644 --- a/pyrit/registry/__init__.py +++ b/pyrit/registry/__init__.py @@ -32,7 +32,6 @@ ) from pyrit.registry.registry import InstanceHoldingRegistry, ParamBagRegistry, Registry from pyrit.registry.registry_metadata import RegistryMetadata - from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage from pyrit.registry.tag_query import TagQuery @@ -54,7 +53,6 @@ "RegistryEntry": "pyrit.registry.instance_registry", "ScenarioMetadata": "pyrit.registry.components", "ScenarioPresetConflictError": "pyrit.registry.scenario_preset_storage", - "ScenarioPresetRegistry": "pyrit.registry.scenario_preset_registry", "ScenarioPresetStorage": "pyrit.registry.scenario_preset_storage", "ScenarioRegistry": "pyrit.registry.components", "ScorerRegistry": "pyrit.registry.components", diff --git a/pyrit/registry/scenario_preset_registry.py b/pyrit/registry/scenario_preset_registry.py deleted file mode 100644 index 174d345760..0000000000 --- a/pyrit/registry/scenario_preset_registry.py +++ /dev/null @@ -1,195 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""Registry unioning built-in and user-authored scenario presets.""" - -from __future__ import annotations - -import logging -from typing import TYPE_CHECKING - -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset -from pyrit.registry.scenario_preset_storage import ScenarioPresetStorage - -if TYPE_CHECKING: - from pathlib import Path - -logger = logging.getLogger(__name__) - - -class ScenarioPresetRegistry: - """ - The union of built-in and user-authored scenario presets. - - Presets arrive from two places. Built-in presets are registered from initializer code - at startup and are read-only, so they are held in memory. User presets are read from - storage on each call rather than cached, because the same directory or blob container - is routinely shared between a notebook, the API, and a second process; a cache would - serve edits those callers can no longer see. The registry exists because only it sees - both sources, so only it can resolve a name across them. - - Built-in presets win on a name collision. A user preset that collides is skipped with a - warning rather than failing startup, because the collision arises when PyRIT ships a new - built-in under a name a user already took, and an upgrade must not break a running - install. To customize a built-in, fork it under a new name. - """ - - def __init__(self, *, storage: ScenarioPresetStorage | None = None) -> None: - """ - Initialize the registry. - - Args: - storage (ScenarioPresetStorage | None): Storage for user presets. Defaults to a - local directory under the PyRIT configuration path. - """ - self._storage = storage - self._builtin_presets: dict[str, ScenarioPreset] = {} - - def configure_storage_source(self, source: str | None) -> None: - """Configure the local directory or Azure Blob source for user presets.""" - self._storage = ScenarioPresetStorage(source=source or str(self._get_default_storage_dir())) - - def register_builtin(self, preset: ScenarioPreset) -> ScenarioPreset: - """ - Register a preset that ships with PyRIT. - - Args: - preset (ScenarioPreset): The preset to register. Its provenance is forced to - built-in so a caller cannot register an editable preset by accident. - - Returns: - ScenarioPreset: The registered preset. - - Raises: - ValueError: If a built-in preset with the same name is already registered. - """ - if preset.name in self._builtin_presets: - raise ValueError(f"Built-in scenario preset '{preset.name}' is already registered.") - - registered = preset.model_copy(update={"provenance": ScenarioPresetProvenance.BUILT_IN}) - self._builtin_presets[registered.name] = registered - return registered - - def get_preset(self, name: str) -> ScenarioPreset | None: - """ - Resolve one preset by name. - - Returns: - ScenarioPreset | None: The preset, or ``None`` if no preset uses that name. - - Raises: - ValueError: If *name* is not a legal preset name. - """ - builtin = self._builtin_presets.get(name) - if builtin is not None: - return builtin - - stored = self._get_storage().load_preset(name) - return stored.preset if stored is not None else None - - def get_stored_preset(self, name: str) -> StoredPreset | None: - """ - Read one user preset together with the version needed to save an edit to it. - - Returns: - StoredPreset | None: The stored preset, or ``None`` if *name* is built in or - no stored preset uses that name. - - Raises: - ValueError: If *name* is not a legal preset name. - """ - if name in self._builtin_presets: - return None - return self._get_storage().load_preset(name) - - def list_presets(self) -> list[ScenarioPreset]: - """ - List every known preset. - - Returns: - list[ScenarioPreset]: Built-in and user presets, sorted by name. - """ - merged: dict[str, ScenarioPreset] = {} - for name, stored in self._get_storage().list_presets().items(): - if name in self._builtin_presets: - logger.warning( - f"Skipping stored scenario preset '{name}': a built-in preset already uses that name. " - "Rename the stored preset to keep using it." - ) - continue - merged[name] = stored.preset - - merged.update(self._builtin_presets) - return [merged[name] for name in sorted(merged)] - - def is_builtin(self, name: str) -> bool: - """Return whether *name* belongs to a built-in preset.""" - return name in self._builtin_presets - - def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) -> StoredPreset: - """ - Persist a user preset. - - Args: - preset (ScenarioPreset): The preset to persist. - expected_version (str | None): ``None`` to create a preset that must not already - exist, or the version returned when the edited preset was read. - - Returns: - StoredPreset: The persisted preset and its new version. - - Raises: - ScenarioPresetConflictError: If the stored version does not match *expected_version*. - ValueError: If the name belongs to a built-in preset or is not a legal preset name. - """ - self._reject_builtin(preset.name, action="saved") - return self._get_storage().save_preset(preset=preset, expected_version=expected_version) - - def delete_preset(self, name: str) -> None: - """ - Delete a user preset from storage. - - Raises: - ValueError: If *name* belongs to a built-in preset or is not a legal preset name. - """ - self._reject_builtin(name, action="deleted") - self._get_storage().delete_preset(name) - - def _reject_builtin(self, name: str, *, action: str) -> None: - """ - Refuse to mutate a built-in preset. - - Raises: - ValueError: If *name* belongs to a built-in preset. - """ - if name in self._builtin_presets: - raise ValueError( - f"Scenario preset '{name}' is built in and cannot be {action}. " - "Fork it under a new name to customize it." - ) - - def _get_storage(self) -> ScenarioPresetStorage: - """ - Return storage for user presets, creating the default backend on first use. - - Returns: - ScenarioPresetStorage: The configured storage backend. - """ - if self._storage is None: - self._storage = ScenarioPresetStorage(source=str(self._get_default_storage_dir())) - return self._storage - - @staticmethod - def _get_default_storage_dir() -> Path: - """ - Get the directory for storing user-authored presets. - - Returns: - Path: Path to ``~/.pyrit/scenario_presets/``, created if needed. - """ - # Deferred: importing pyrit.common.path triggers pyrit __init__.py - from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH - - presets_dir = CONFIGURATION_DIRECTORY_PATH / "scenario_presets" - presets_dir.mkdir(parents=True, exist_ok=True) - return presets_dir diff --git a/pyrit/registry/scenario_preset_storage.py b/pyrit/registry/scenario_preset_storage.py index bce730045f..b94a86a680 100644 --- a/pyrit/registry/scenario_preset_storage.py +++ b/pyrit/registry/scenario_preset_storage.py @@ -1,17 +1,21 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -"""Storage backends for user-authored scenario presets.""" +"""Storage backends for scenario presets.""" from __future__ import annotations import hashlib import json import logging +from typing import TYPE_CHECKING -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance, StoredPreset +from pyrit.models.catalog.scenario_preset import ScenarioPreset, StoredPreset from pyrit.registry.file_document_storage import FileDocumentStorage +if TYPE_CHECKING: + from pathlib import Path + logger = logging.getLogger(__name__) @@ -44,10 +48,11 @@ def _document_version(content: str) -> str: class ScenarioPresetStorage(FileDocumentStorage): """ - Read and write user-authored scenario presets as JSON documents. + Read and write scenario presets as JSON documents. - Only user presets are stored. Built-in presets come from initializer code and are - never written here, so a file on disk always means a user authored it. + Storage is read through on every call rather than cached, because the same directory + or blob container is routinely shared between a notebook, the API, and a second + process; a cache would serve edits those callers can no longer see. Writes are guarded by an optimistic-concurrency check. A caller supplies the version it read, and the save is refused unless storage still holds that version. The token @@ -61,14 +66,37 @@ class ScenarioPresetStorage(FileDocumentStorage): writes (blob ETags have no local-filesystem equivalent) and is deliberately deferred. """ - def __init__(self, *, source: str) -> None: + def __init__(self, *, source: str | None = None) -> None: """ Initialize storage from a local directory or Azure Blob source URI. + Args: + source (str | None): Local directory or Azure Blob source URI. Defaults to + ``scenario_presets`` under the PyRIT configuration directory. + Raises: ValueError: If the source has an unsupported URI scheme. """ - super().__init__(source=source, extension=".json", source_label="Scenario preset") + super().__init__( + source=source or str(self._get_default_storage_dir()), + extension=".json", + source_label="Scenario preset", + ) + + @staticmethod + def _get_default_storage_dir() -> Path: + """ + Get the default directory for storing presets. + + Returns: + Path: Path to ``~/.pyrit/scenario_presets/``, created if needed. + """ + # Deferred: importing pyrit.common.path triggers pyrit __init__.py + from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH + + presets_dir = CONFIGURATION_DIRECTORY_PATH / "scenario_presets" + presets_dir.mkdir(parents=True, exist_ok=True) + return presets_dir def get_preset_source(self, name: str) -> str: """ @@ -148,10 +176,9 @@ def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) - name=preset.name, expected_version=expected_version, actual_version=actual_version ) - saved = preset.model_copy(update={"provenance": ScenarioPresetProvenance.USER}) - content = self._serialize_preset(saved) - self._save_document(name=saved.name, content=content) - return StoredPreset(preset=saved, version=_document_version(content)) + content = self._serialize_preset(preset) + self._save_document(name=preset.name, content=content) + return StoredPreset(preset=preset, version=_document_version(content)) def delete_preset(self, name: str) -> None: """ @@ -197,7 +224,7 @@ def _parse_preset(*, name: str, content: str) -> ScenarioPreset | None: logger.error(f"Skipping stored scenario preset '{name}': it is not a JSON object.") return None - fields = {**payload, "name": name, "provenance": ScenarioPresetProvenance.USER} + fields = {**payload, "name": name} try: return ScenarioPreset.model_validate(fields) except Exception: diff --git a/tests/unit/models/test_scenario_preset.py b/tests/unit/models/test_scenario_preset.py index 5e8513ab4f..7493139409 100644 --- a/tests/unit/models/test_scenario_preset.py +++ b/tests/unit/models/test_scenario_preset.py @@ -5,7 +5,7 @@ import pytest -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.models.catalog.scenario_preset import ScenarioPreset def test_init_defaults_every_optional_field_to_none() -> None: @@ -21,30 +21,11 @@ def test_init_defaults_every_optional_field_to_none() -> None: assert preset.description is None -def test_init_defaults_to_user_provenance() -> None: - """Test the default identity of a freshly constructed preset.""" - preset = ScenarioPreset(name="nightly", scenario_name="foundry.red_team_agent") - - assert preset.provenance is ScenarioPresetProvenance.USER - assert preset.is_builtin is False - - def test_preset_carries_no_storage_version() -> None: """Test that storage metadata stays out of the model it describes.""" assert "version" not in ScenarioPreset.model_fields -def test_is_builtin_reflects_provenance() -> None: - """Test that built-in provenance marks a preset read-only.""" - preset = ScenarioPreset( - name="nightly", - scenario_name="foundry.red_team_agent", - provenance=ScenarioPresetProvenance.BUILT_IN, - ) - - assert preset.is_builtin is True - - @pytest.mark.parametrize("name", ["Nightly", "nightly-scan", "1nightly", "", "a" * 65]) def test_init_rejects_invalid_registry_names(name: str) -> None: """Test that a preset name must be a legal registry name.""" diff --git a/tests/unit/registry/test_scenario_preset_registry.py b/tests/unit/registry/test_scenario_preset_registry.py deleted file mode 100644 index bf96bd645c..0000000000 --- a/tests/unit/registry/test_scenario_preset_registry.py +++ /dev/null @@ -1,251 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""Tests for the scenario preset registry.""" - -import logging -from pathlib import Path - -import pytest - -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance -from pyrit.registry.scenario_preset_registry import ScenarioPresetRegistry -from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage - - -def _make_preset(**overrides: object) -> ScenarioPreset: - """ - Build a preset with test defaults. - - Returns: - ScenarioPreset: The constructed preset. - """ - fields: dict[str, object] = {"name": "nightly", "scenario_name": "foundry.red_team_agent"} - fields.update(overrides) - return ScenarioPreset(**fields) # type: ignore[arg-type] - - -@pytest.fixture -def registry(tmp_path: Path) -> ScenarioPresetRegistry: - """ - Build a registry backed by an isolated local directory. - - Returns: - ScenarioPresetRegistry: The registry under test. - """ - return ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(tmp_path))) - - -def test_list_presets_returns_union_sorted_by_name(registry: ScenarioPresetRegistry) -> None: - """Test that the registry exposes code-registered and file-loaded presets together.""" - registry.register_builtin(_make_preset(name="builtin_suite")) - registry.save_preset(preset=_make_preset(name="user_suite"), expected_version=None) - - names = [preset.name for preset in registry.list_presets()] - - assert names == ["builtin_suite", "user_suite"] - - -def test_register_builtin_forces_builtin_provenance(registry: ScenarioPresetRegistry) -> None: - """Test that a registered built-in cannot be marked editable by its author.""" - registered = registry.register_builtin(_make_preset(provenance=ScenarioPresetProvenance.USER)) - - assert registered.provenance is ScenarioPresetProvenance.BUILT_IN - assert registry.is_builtin("nightly") is True - - -def test_register_builtin_rejects_duplicate_name(registry: ScenarioPresetRegistry) -> None: - """Test that two built-ins cannot claim the same name.""" - registry.register_builtin(_make_preset()) - - with pytest.raises(ValueError, match="already registered"): - registry.register_builtin(_make_preset()) - - -def test_get_preset_resolves_both_provenances(registry: ScenarioPresetRegistry) -> None: - """Test name resolution across both sources.""" - registry.register_builtin(_make_preset(name="builtin_suite")) - registry.save_preset(preset=_make_preset(name="user_suite"), expected_version=None) - - assert registry.get_preset("builtin_suite") is not None - assert registry.get_preset("user_suite") is not None - assert registry.get_preset("absent") is None - - -def test_get_preset_rejects_illegal_names(registry: ScenarioPresetRegistry) -> None: - """Test that a traversal attempt fails loudly rather than resolving to nothing.""" - with pytest.raises(ValueError, match="Invalid registry name"): - registry.get_preset("../victim") - - -def test_get_stored_preset_returns_the_version_needed_to_edit(registry: ScenarioPresetRegistry) -> None: - """Test that an editor can read a preset and the token required to save it back.""" - saved = registry.save_preset(preset=_make_preset(), expected_version=None) - - stored = registry.get_stored_preset("nightly") - - assert stored is not None - assert stored.version == saved.version - registry.save_preset(preset=_make_preset(description="edited"), expected_version=stored.version) - - -def test_get_stored_preset_returns_none_for_a_builtin(registry: ScenarioPresetRegistry) -> None: - """Test that a built-in has no stored document and therefore no editable version.""" - registry.register_builtin(_make_preset()) - - assert registry.get_stored_preset("nightly") is None - - -def test_saving_over_a_builtin_name_is_rejected(registry: ScenarioPresetRegistry) -> None: - """Test that forking a built-in requires a new name.""" - registry.register_builtin(_make_preset()) - - with pytest.raises(ValueError, match="built in and cannot be saved"): - registry.save_preset(preset=_make_preset(), expected_version=None) - - -def test_deleting_a_builtin_is_rejected(registry: ScenarioPresetRegistry) -> None: - """Test that built-in presets cannot be removed.""" - registry.register_builtin(_make_preset()) - - with pytest.raises(ValueError, match="built in and cannot be deleted"): - registry.delete_preset("nightly") - - -def test_fork_under_a_new_name_succeeds(registry: ScenarioPresetRegistry) -> None: - """Test that the supported customization path works.""" - registry.register_builtin(_make_preset(name="builtin_suite", max_dataset_size=200)) - - forked = registry.save_preset( - preset=_make_preset(name="builtin_suite_fork", max_dataset_size=20), expected_version=None - ) - - assert forked.preset.name == "builtin_suite_fork" - assert forked.preset.max_dataset_size == 20 - assert registry.get_preset("builtin_suite") is not None - - -def test_builtin_wins_collision_and_user_preset_is_skipped_with_warning( - tmp_path: Path, caplog: pytest.LogCaptureFixture -) -> None: - """Test the upgrade case: a new built-in shadows a stored preset without failing startup.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - storage.save_preset(preset=_make_preset(description="user version"), expected_version=None) - - registry = ScenarioPresetRegistry(storage=storage) - registry.register_builtin(_make_preset(description="shipped version")) - - with caplog.at_level(logging.WARNING): - listed = registry.list_presets() - - resolved = registry.get_preset("nightly") - assert resolved is not None - assert resolved.description == "shipped version" - assert resolved.is_builtin is True - assert [preset.name for preset in listed] == ["nightly"] - assert "a built-in preset already uses that name" in caplog.text - - -def test_collision_skips_only_the_colliding_preset(tmp_path: Path) -> None: - """Test that a collision skips only the colliding preset.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - storage.save_preset(preset=_make_preset(name="nightly"), expected_version=None) - storage.save_preset(preset=_make_preset(name="weekly"), expected_version=None) - - registry = ScenarioPresetRegistry(storage=storage) - registry.register_builtin(_make_preset(name="nightly")) - - assert [preset.name for preset in registry.list_presets()] == ["nightly", "weekly"] - weekly = registry.get_preset("weekly") - assert weekly is not None - assert weekly.is_builtin is False - - -def test_presets_written_by_another_process_are_visible(tmp_path: Path) -> None: - """Test that user presets are read through, so a shared source is not served from a stale cache.""" - registry = ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(tmp_path))) - assert registry.get_preset("nightly") is None - - ScenarioPresetStorage(source=str(tmp_path)).save_preset(preset=_make_preset(), expected_version=None) - - assert registry.get_preset("nightly") is not None - assert [preset.name for preset in registry.list_presets()] == ["nightly"] - - -def test_presets_deleted_by_another_process_disappear(tmp_path: Path) -> None: - """Test that a preset removed outside the registry stops resolving.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - registry = ScenarioPresetRegistry(storage=storage) - registry.save_preset(preset=_make_preset(), expected_version=None) - - storage.delete_preset("nightly") - - assert registry.get_preset("nightly") is None - - -def test_builtin_presets_are_never_written_to_storage(tmp_path: Path) -> None: - """Test that registering a built-in does not persist anything.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - registry = ScenarioPresetRegistry(storage=storage) - - registry.register_builtin(_make_preset()) - - assert storage.list_presets() == {} - assert list(tmp_path.glob("*.json")) == [] - - -def test_save_preset_is_immediately_visible(registry: ScenarioPresetRegistry) -> None: - """Test that a save is visible through the registry without reloading.""" - first = registry.save_preset(preset=_make_preset(description="first"), expected_version=None) - registry.save_preset(preset=_make_preset(description="second"), expected_version=first.version) - - resolved = registry.get_preset("nightly") - - assert resolved is not None - assert resolved.description == "second" - - -def test_stale_save_propagates_conflict(registry: ScenarioPresetRegistry) -> None: - """Test that the registry surfaces the storage concurrency check.""" - registry.save_preset(preset=_make_preset(), expected_version=None) - - with pytest.raises(ScenarioPresetConflictError): - registry.save_preset(preset=_make_preset(description="stale"), expected_version="stale_version") - - -def test_delete_preset_removes_from_registry_and_storage(tmp_path: Path) -> None: - """Test that deletion clears the backing file and stops resolution.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) - registry = ScenarioPresetRegistry(storage=storage) - registry.save_preset(preset=_make_preset(), expected_version=None) - - registry.delete_preset("nightly") - - assert registry.get_preset("nightly") is None - assert storage.list_presets() == {} - - -def test_configure_storage_source_switches_backend(tmp_path: Path) -> None: - """Test that the storage source can be pointed at an explicit directory.""" - registry = ScenarioPresetRegistry() - presets_dir = tmp_path / "presets" - presets_dir.mkdir() - - registry.configure_storage_source(str(presets_dir)) - registry.save_preset(preset=_make_preset(), expected_version=None) - - assert (presets_dir / "nightly.json").is_file() - - -def test_configure_storage_source_stops_serving_the_previous_source(tmp_path: Path) -> None: - """Test that repointing the source does not leave presets from the old one resolvable.""" - first_dir = tmp_path / "first" - second_dir = tmp_path / "second" - first_dir.mkdir() - second_dir.mkdir() - registry = ScenarioPresetRegistry(storage=ScenarioPresetStorage(source=str(first_dir))) - registry.save_preset(preset=_make_preset(), expected_version=None) - - registry.configure_storage_source(str(second_dir)) - - assert registry.get_preset("nightly") is None diff --git a/tests/unit/registry/test_scenario_preset_storage.py b/tests/unit/registry/test_scenario_preset_storage.py index 9b8a061b94..ee679397ec 100644 --- a/tests/unit/registry/test_scenario_preset_storage.py +++ b/tests/unit/registry/test_scenario_preset_storage.py @@ -10,7 +10,7 @@ import pytest -from pyrit.models.catalog.scenario_preset import ScenarioPreset, ScenarioPresetProvenance +from pyrit.models.catalog.scenario_preset import ScenarioPreset from pyrit.registry.scenario_preset_storage import ScenarioPresetConflictError, ScenarioPresetStorage @@ -189,18 +189,34 @@ def test_update_of_missing_preset_is_rejected(tmp_path: Path) -> None: assert error.value.actual_version is None -def test_save_forces_user_provenance(tmp_path: Path) -> None: - """Test that stored presets are always user-owned regardless of the submitted value.""" - storage = ScenarioPresetStorage(source=str(tmp_path)) +def test_default_source_is_under_the_configuration_directory(tmp_path: Path) -> None: + """Test that omitting the source stores presets under the PyRIT configuration directory.""" + with patch("pyrit.common.path.CONFIGURATION_DIRECTORY_PATH", tmp_path): + storage = ScenarioPresetStorage() + storage.save_preset(preset=_make_preset(), expected_version=None) - saved = storage.save_preset( - preset=_make_preset(provenance=ScenarioPresetProvenance.BUILT_IN), expected_version=None - ) + assert (tmp_path / "scenario_presets" / "nightly.json").is_file() - assert saved.preset.provenance is ScenarioPresetProvenance.USER - loaded = storage.load_preset("nightly") - assert loaded is not None - assert loaded.preset.provenance is ScenarioPresetProvenance.USER + +def test_presets_written_by_another_process_are_visible(tmp_path: Path) -> None: + """Test that reads go through to the source, so a shared directory is never served stale.""" + reader = ScenarioPresetStorage(source=str(tmp_path)) + assert reader.load_preset("nightly") is None + + ScenarioPresetStorage(source=str(tmp_path)).save_preset(preset=_make_preset(), expected_version=None) + + assert reader.load_preset("nightly") is not None + assert list(reader.list_presets()) == ["nightly"] + + +def test_presets_deleted_by_another_process_disappear(tmp_path: Path) -> None: + """Test that a preset removed outside this instance stops resolving.""" + reader = ScenarioPresetStorage(source=str(tmp_path)) + reader.save_preset(preset=_make_preset(), expected_version=None) + + ScenarioPresetStorage(source=str(tmp_path)).delete_preset("nightly") + + assert reader.load_preset("nightly") is None def test_load_missing_preset_returns_none(tmp_path: Path) -> None: From 246a7c791548c3a1b10a3cf0578f92672a20d763 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Thu, 1 Oct 2026 18:31:07 -0400 Subject: [PATCH 04/16] FIX: Skip unaddressable stored document names and harden preset documents Listing no longer returns names that the single-document operations reject. FileDocumentStorage._list_documents now filters enumerated names through validate_registry_name and logs a warning for the ones it skips, so list_scripts() and list_presets() only return names that get/save/delete will accept. This fixes a 500 from the custom initializer list endpoint whenever the script directory held an ordinary file like __init__.py or My-Script.py, and a misleading 400 'Invalid registry name' from the delete endpoint for a file that demonstrably exists. Also: ScenarioPreset now forbids unknown fields so a misspelled key fails loudly instead of being silently dropped; the preset name is no longer duplicated inside the document, since the document name is authoritative; get_preset_version() lets a caller recover a preset name whose document is malformed; and RunScenarioRequest.max_dataset_size no longer describes the removed per-dataset behavior. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/models/catalog/scenario.py | 2 +- pyrit/models/catalog/scenario_preset.py | 9 +++- pyrit/registry/file_document_storage.py | 39 ++++++++++++-- pyrit/registry/scenario_preset_storage.py | 25 ++++++++- tests/unit/models/test_scenario_preset.py | 10 ++++ .../test_custom_initializer_storage.py | 13 +++++ .../registry/test_initializer_registry.py | 15 ++++++ .../registry/test_scenario_preset_storage.py | 54 +++++++++++++++++++ 8 files changed, 160 insertions(+), 7 deletions(-) diff --git a/pyrit/models/catalog/scenario.py b/pyrit/models/catalog/scenario.py index bd8d07230a..7692580eca 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -393,7 +393,7 @@ class RunScenarioRequest(BaseModel): ) techniques: list[str] | None = Field(None, description="Technique names to use (uses scenario default if omitted)") dataset_names: list[str] | None = Field(None, description="Dataset names to use (uses scenario default if omitted)") - max_dataset_size: int | None = Field(None, ge=1, description="Maximum items per dataset") + max_dataset_size: int | None = Field(None, ge=1, description="Maximum selected logical seed groups") dataset_filters: dict[str, list[str]] | None = Field( None, description=( diff --git a/pyrit/models/catalog/scenario_preset.py b/pyrit/models/catalog/scenario_preset.py index 4583e5a94e..2baaf29231 100644 --- a/pyrit/models/catalog/scenario_preset.py +++ b/pyrit/models/catalog/scenario_preset.py @@ -18,7 +18,7 @@ from typing import Any -from pydantic import BaseModel, Field, field_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator from pyrit.models.catalog.scenario import _validate_dataset_filter_mapping from pyrit.models.identifiers.class_name_utils import validate_registry_name @@ -31,8 +31,15 @@ class ScenarioPreset(BaseModel): Presets own *what to test*. A launch owns *how and where* — the target, concurrency, retries, and labels — so those fields are absent here by design and the two sets are combined by union rather than by precedence. + + Unknown keys are rejected rather than ignored. These documents are hand-edited, and + because an absent field is a meaningful state, a misspelled key would otherwise parse + cleanly and leave the preset silently testing the scenario default instead of the + value the file plainly states. """ + model_config = ConfigDict(extra="forbid") + name: str = Field(..., description="Unique preset name, used as the storage key and as the reference from scans") scenario_name: str = Field(..., min_length=1, description="Registered scenario this preset configures") description: str | None = Field(None, description="Human-readable summary of what this preset tests") diff --git a/pyrit/registry/file_document_storage.py b/pyrit/registry/file_document_storage.py index c6cd56e57d..bbf6c57f5b 100644 --- a/pyrit/registry/file_document_storage.py +++ b/pyrit/registry/file_document_storage.py @@ -5,6 +5,7 @@ from __future__ import annotations +import logging from contextlib import contextmanager, suppress from pathlib import Path, PurePosixPath from typing import TYPE_CHECKING @@ -18,6 +19,8 @@ from azure.storage.blob import ContainerClient +logger = logging.getLogger(__name__) + class FileDocumentStorage: """ @@ -32,6 +35,11 @@ class FileDocumentStorage: a caller read, overwrite, or delete a file outside the configured source. The check lives here rather than in each subclass so no document API can omit it. + Listing applies the same rule, so every name it returns can be passed back to the + single-document operations. Names come from the file system rather than from a + caller, so the source can hold files this class cannot address; those are skipped + with a warning rather than failing the whole listing. + Subclasses supply the extension and a human-readable label for error messages, then expose a domain-specific API over the protected document operations. """ @@ -84,17 +92,42 @@ def _get_document_source(self, name: str) -> str: def _list_documents(self) -> dict[str, str]: """ - Read every stored document. + Read every stored document that the single-document operations can address. Returns: dict[str, str]: Document content keyed by name. """ - if self._is_blob: - return self._list_blob_documents() + documents = self._list_blob_documents() if self._is_blob else self._list_local_documents() + return {name: content for name, content in documents.items() if self._is_addressable_name(name)} + def _list_local_documents(self) -> dict[str, str]: + """ + Read documents from the configured local directory. + + Returns: + dict[str, str]: Document content keyed by file stem. + """ directory = self._local_directory(create=True) return {path.stem: path.read_text(encoding="utf-8") for path in sorted(directory.glob(f"*{self._extension}"))} + def _is_addressable_name(self, name: str) -> bool: + """ + Return whether a discovered document name is one this storage can address. + + Ordinary files such as ``__init__.py`` or ``My-Script.py`` can sit alongside + valid documents, and a caller that fed such a name back into a read, write, or + delete would get a ``ValueError`` it has no way to anticipate. + + Returns: + bool: Whether the name is a legal registry name. + """ + try: + validate_registry_name(name) + except ValueError as error: + logger.warning(f"Ignoring stored document '{name}{self._extension}' in {self.display_source}: {error}") + return False + return True + def _read_document(self, name: str) -> str | None: """ Read one document. diff --git a/pyrit/registry/scenario_preset_storage.py b/pyrit/registry/scenario_preset_storage.py index b94a86a680..b52304303e 100644 --- a/pyrit/registry/scenario_preset_storage.py +++ b/pyrit/registry/scenario_preset_storage.py @@ -146,6 +146,24 @@ def load_preset(self, name: str) -> StoredPreset | None: return None return StoredPreset(preset=preset, version=_document_version(content)) + def get_preset_version(self, name: str) -> str | None: + """ + Read the version of one stored document without parsing it. + + A document that cannot be parsed is otherwise unreachable: ``list_presets`` skips + it, ``load_preset`` returns ``None``, and a create is refused because the document + exists. Exposing its version lets a caller offer to overwrite the broken file + instead of leaving the name permanently unusable. + + Returns: + str | None: The document version, or ``None`` if no document is stored. + + Raises: + ValueError: If *name* is not a legal preset name. + """ + content = self._read_document(name) + return None if content is None else _document_version(content) + def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) -> StoredPreset: """ Persist one preset. @@ -195,12 +213,15 @@ def _serialize_preset(preset: ScenarioPreset) -> str: Serialize a preset to stored JSON. Unset fields are omitted rather than written as ``null`` so a stored preset reads - as the set of decisions its author actually made. + as the set of decisions its author actually made. The name is omitted too: it is + the document key, and writing it would invite a hand-editor to change it and + expect a rename that cannot happen. Returns: str: JSON document content. """ - return json.dumps(preset.model_dump(mode="json", exclude_none=True), indent=2, sort_keys=True) + "\n" + payload = preset.model_dump(mode="json", exclude_none=True, exclude={"name"}) + return json.dumps(payload, indent=2, sort_keys=True) + "\n" @staticmethod def _parse_preset(*, name: str, content: str) -> ScenarioPreset | None: diff --git a/tests/unit/models/test_scenario_preset.py b/tests/unit/models/test_scenario_preset.py index 7493139409..588adbf366 100644 --- a/tests/unit/models/test_scenario_preset.py +++ b/tests/unit/models/test_scenario_preset.py @@ -26,6 +26,16 @@ def test_preset_carries_no_storage_version() -> None: assert "version" not in ScenarioPreset.model_fields +def test_init_rejects_unknown_field() -> None: + """Test that a misspelled key fails loudly instead of silently falling back to the default.""" + with pytest.raises(ValueError, match="techinques"): + ScenarioPreset( + name="nightly", + scenario_name="foundry.red_team_agent", + techinques=["crescendo"], # type: ignore[call-arg] + ) + + @pytest.mark.parametrize("name", ["Nightly", "nightly-scan", "1nightly", "", "a" * 65]) def test_init_rejects_invalid_registry_names(name: str) -> None: """Test that a preset name must be a legal registry name.""" diff --git a/tests/unit/registry/test_custom_initializer_storage.py b/tests/unit/registry/test_custom_initializer_storage.py index cace9d16a2..3da0541888 100644 --- a/tests/unit/registry/test_custom_initializer_storage.py +++ b/tests/unit/registry/test_custom_initializer_storage.py @@ -135,6 +135,19 @@ def test_local_storage_reads_latest_script_content(tmp_path: Path) -> None: assert storage.list_scripts() == {"example": "VALUE = 2\n"} +def test_listing_skips_scripts_it_cannot_address(tmp_path: Path) -> None: + """Test that every listed name can be passed back into the single-document operations.""" + for file_name in ["good_one.py", "My-Script.py", "__init__.py", "test-helper.py"]: + (tmp_path / file_name).write_text("VALUE = 1\n", encoding="utf-8") + storage = CustomInitializerStorage(source=str(tmp_path)) + + listed = storage.list_scripts() + + assert sorted(listed) == ["good_one"] + for name in listed: + storage.get_script_source(name) + + def test_direct_python_blob_rejects_name_outside_prefix(tmp_path: Path) -> None: """Test that blobs outside the configured virtual directory are ignored.""" storage = CustomInitializerStorage(source=str(tmp_path)) diff --git a/tests/unit/registry/test_initializer_registry.py b/tests/unit/registry/test_initializer_registry.py index a4991c377a..9aae284e3e 100644 --- a/tests/unit/registry/test_initializer_registry.py +++ b/tests/unit/registry/test_initializer_registry.py @@ -265,6 +265,21 @@ def test_list_stored_initializer_sources_includes_display_paths(lazy_registry: I ) +def test_list_stored_initializer_sources_tolerates_unaddressable_files( + lazy_registry: InitializerRegistry, tmp_path: Path +) -> None: + """Test that a stray file in the source directory does not fail the whole listing.""" + (tmp_path / "good_one.py").write_text(_VALID_SCRIPT, encoding="utf-8") + (tmp_path / "__init__.py").write_text(_VALID_SCRIPT, encoding="utf-8") + (tmp_path / "My-Script.py").write_text(_VALID_SCRIPT, encoding="utf-8") + lazy_registry.configure_custom_scripts_source(str(tmp_path)) + + source, items = lazy_registry.list_stored_initializer_sources() + + assert source == str(tmp_path) + assert [name for name, _, _ in items] == ["good_one"] + + def test_unregister_and_cleanup_rejects_builtin(lazy_registry): """Test that unregister_and_cleanup raises ValueError for built-in initializers.""" diff --git a/tests/unit/registry/test_scenario_preset_storage.py b/tests/unit/registry/test_scenario_preset_storage.py index ee679397ec..8bf4a140b5 100644 --- a/tests/unit/registry/test_scenario_preset_storage.py +++ b/tests/unit/registry/test_scenario_preset_storage.py @@ -350,3 +350,57 @@ def test_blob_storage_reads_missing_document_as_none() -> None: with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): assert storage.load_preset("absent") is None + + +def test_name_is_not_written_into_the_document(tmp_path: Path) -> None: + """Test that the storage key appears only in the file name, so no hand-edit can contradict it.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + storage.save_preset(preset=_make_preset(), expected_version=None) + + assert "name" not in json.loads((tmp_path / "nightly.json").read_text(encoding="utf-8")) + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.preset.name == "nightly" + + +def test_misspelled_key_is_skipped_not_silently_dropped(tmp_path: Path) -> None: + """Test that a typo in a hand-edited file is refused rather than quietly using the scenario default.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + (tmp_path / "nightly.json").write_text( + json.dumps({"scenario_name": "foundry.red_team_agent", "techinques": ["crescendo"]}), + encoding="utf-8", + ) + + assert storage.load_preset("nightly") is None + assert storage.list_presets() == {} + + +def test_malformed_document_can_be_overwritten_through_its_version(tmp_path: Path) -> None: + """Test that a file which cannot be parsed still has a recovery path rather than burning the name.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + (tmp_path / "nightly.json").write_text("{not json", encoding="utf-8") + + version = storage.get_preset_version("nightly") + assert version is not None + + saved = storage.save_preset(preset=_make_preset(), expected_version=version) + + assert saved.preset.name == "nightly" + assert storage.load_preset("nightly") is not None + + +def test_get_preset_version_returns_none_for_absent_document(tmp_path: Path) -> None: + """Test that an absent document reports no version, so a create still reads as a create.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + assert storage.get_preset_version("nightly") is None + + +def test_listing_skips_documents_it_cannot_address(tmp_path: Path) -> None: + """Test that a stray file whose stem is not a legal name does not fail the whole listing.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(), expected_version=None) + (tmp_path / "My-Preset.json").write_text(json.dumps({"scenario_name": "foundry.red_team_agent"}), encoding="utf-8") + + assert sorted(storage.list_presets()) == ["nightly"] From 3086224e1c241450dbd24ed7585ee2848a4b1c80 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Fri, 2 Oct 2026 13:04:51 -0400 Subject: [PATCH 05/16] FIX: Read stored documents as bytes and replace them atomically Listing read and decoded every matching file before checking whether the name was one it could address, so a single undecodable file raised out of list_presets and list_scripts and hid every other stored document. Each backend now validates the name first, reads per document, and skips one it cannot read with a warning. Decoding moves out of storage into the subclasses, which know how to report a document they cannot interpret. That also makes the version token hash the exact bytes on disk. Text mode had been rewriting line endings on Windows, so the token described the in-memory string rather than the stored file, and a raw-byte read for version checks alone would have made every update conflict. Local writes stage the content in a sibling temporary file and move it into place, so a failed write leaves the previous document intact instead of truncating it, and a concurrent reader never sees a partial document. Initializer source is normalized on decode. Callers materialize it back to a file in text mode, which would otherwise translate a stored CRLF into CRCRLF. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/registry/custom_initializer_storage.py | 33 ++++- pyrit/registry/file_document_storage.py | 108 ++++++++++++---- pyrit/registry/scenario_preset_storage.py | 33 +++-- .../test_custom_initializer_storage.py | 31 +++++ .../registry/test_scenario_preset_storage.py | 120 ++++++++++++++++++ 5 files changed, 283 insertions(+), 42 deletions(-) diff --git a/pyrit/registry/custom_initializer_storage.py b/pyrit/registry/custom_initializer_storage.py index 8b38323312..572c838c00 100644 --- a/pyrit/registry/custom_initializer_storage.py +++ b/pyrit/registry/custom_initializer_storage.py @@ -5,8 +5,28 @@ from __future__ import annotations +import logging + from pyrit.registry.file_document_storage import FileDocumentStorage +logger = logging.getLogger(__name__) + + +def _decode_script(content: bytes) -> str: + """ + Decode stored script bytes with the line endings a text-mode read would produce. + + Callers materialize this source back to a file in text mode, which translates every + newline again, so a stored CRLF would round-trip into a corrupted CRCRLF. + + Returns: + str: Decoded script source with normalized line endings. + + Raises: + UnicodeDecodeError: If the content is not valid UTF-8 text. + """ + return content.decode("utf-8").replace("\r\n", "\n").replace("\r", "\n") + class CustomInitializerStorage(FileDocumentStorage): """Read and write custom initializer scripts in a directory or blob container.""" @@ -33,14 +53,23 @@ def list_scripts(self) -> dict[str, str]: """ List stored Python scripts by registry name. + A script that is not valid UTF-8 text is logged and skipped, so one unreadable + file cannot hide every other stored initializer. + Returns: dict[str, str]: Script content keyed by registry name. """ - return self._list_documents() + scripts: dict[str, str] = {} + for name, content in self._list_documents().items(): + try: + scripts[name] = _decode_script(content) + except UnicodeDecodeError: + logger.warning(f"Skipping stored initializer '{name}': it is not valid UTF-8 text.") + return scripts def save_script(self, *, name: str, content: str) -> None: """Persist one custom initializer script.""" - self._save_document(name=name, content=content) + self._save_document(name=name, content=content.encode("utf-8")) def delete_script(self, name: str) -> None: """Delete one custom initializer script if it exists.""" diff --git a/pyrit/registry/file_document_storage.py b/pyrit/registry/file_document_storage.py index bbf6c57f5b..6d9e1cf71c 100644 --- a/pyrit/registry/file_document_storage.py +++ b/pyrit/registry/file_document_storage.py @@ -6,6 +6,8 @@ from __future__ import annotations import logging +import os +import tempfile from contextlib import contextmanager, suppress from pathlib import Path, PurePosixPath from typing import TYPE_CHECKING @@ -35,10 +37,18 @@ class FileDocumentStorage: a caller read, overwrite, or delete a file outside the configured source. The check lives here rather than in each subclass so no document API can omit it. - Listing applies the same rule, so every name it returns can be passed back to the - single-document operations. Names come from the file system rather than from a - caller, so the source can hold files this class cannot address; those are skipped - with a warning rather than failing the whole listing. + Listing applies the same rule before it reads anything, so every name it returns can + be passed back to the single-document operations. Names come from the file system + rather than from a caller, so the source can hold files this class cannot address or + cannot read; those are skipped with a warning rather than failing the whole listing. + + Documents are read and written as raw bytes. Decoding belongs to the subclasses, + which know how to report a document they cannot interpret and can skip just that one. + Bytes also keep what a caller hashes identical to what is stored, which text mode + would not: it rewrites line endings per platform. + + Local writes stage the content beside the destination and move it into place, so an + interrupted write leaves the previous document intact rather than truncating it. Subclasses supply the extension and a human-readable label for error messages, then expose a domain-specific API over the protected document operations. @@ -90,25 +100,35 @@ def _get_document_source(self, name: str) -> str: return f"{self.display_source.rstrip('/')}/{name}{self._extension}" return str(self._local_directory() / f"{name}{self._extension}") - def _list_documents(self) -> dict[str, str]: + def _list_documents(self) -> dict[str, bytes]: """ Read every stored document that the single-document operations can address. + A name is checked before its content is read, so a file this class cannot address + is never opened and cannot fail the listing on its way out. + Returns: - dict[str, str]: Document content keyed by name. + dict[str, bytes]: Document content keyed by name. """ - documents = self._list_blob_documents() if self._is_blob else self._list_local_documents() - return {name: content for name, content in documents.items() if self._is_addressable_name(name)} + return self._list_blob_documents() if self._is_blob else self._list_local_documents() - def _list_local_documents(self) -> dict[str, str]: + def _list_local_documents(self) -> dict[str, bytes]: """ - Read documents from the configured local directory. + Read addressable documents from the configured local directory. Returns: - dict[str, str]: Document content keyed by file stem. + dict[str, bytes]: Document content keyed by file stem. """ directory = self._local_directory(create=True) - return {path.stem: path.read_text(encoding="utf-8") for path in sorted(directory.glob(f"*{self._extension}"))} + documents: dict[str, bytes] = {} + for path in sorted(directory.glob(f"*{self._extension}")): + if not self._is_addressable_name(path.stem): + continue + try: + documents[path.stem] = path.read_bytes() + except OSError as error: + logger.warning(f"Skipping unreadable document '{path.name}' in {self.display_source}: {error}") + return documents def _is_addressable_name(self, name: str) -> bool: """ @@ -128,12 +148,12 @@ def _is_addressable_name(self, name: str) -> bool: return False return True - def _read_document(self, name: str) -> str | None: + def _read_document_bytes(self, name: str) -> bytes | None: """ - Read one document. + Read the raw bytes of one document. Returns: - str | None: Document content, or ``None`` if it does not exist. + bytes | None: Document content, or ``None`` if it does not exist. Raises: ValueError: If *name* is not a legal registry name. @@ -144,16 +164,16 @@ def _read_document(self, name: str) -> str | None: with self._open_container_client() as client: try: - return client.download_blob(self._get_blob_name(name)).readall().decode("utf-8") + return client.download_blob(self._get_blob_name(name)).readall() except ResourceNotFoundError: return None path = self._local_directory() / f"{name}{self._extension}" - return path.read_text(encoding="utf-8") if path.is_file() else None + return path.read_bytes() if path.is_file() else None - def _save_document(self, *, name: str, content: str) -> None: + def _save_document(self, *, name: str, content: bytes) -> None: """ - Persist one document, overwriting any existing content. + Persist one document, replacing any existing content. Raises: ValueError: If *name* is not a legal registry name. @@ -161,10 +181,37 @@ def _save_document(self, *, name: str, content: str) -> None: validate_registry_name(name) if self._is_blob: with self._open_container_client() as client: - client.upload_blob(name=self._get_blob_name(name), data=content.encode("utf-8"), overwrite=True) + client.upload_blob(name=self._get_blob_name(name), data=content, overwrite=True) else: directory = self._local_directory(create=True) - (directory / f"{name}{self._extension}").write_text(content, encoding="utf-8") + self._replace_file(path=directory / f"{name}{self._extension}", content=content) + + @staticmethod + def _replace_file(*, path: Path, content: bytes) -> None: + """ + Write *content* to *path* without destroying what is already there on failure. + + Writing in place truncates the destination before the new content lands, so an + interrupted write would leave the stored document empty and a concurrent reader + could observe a half-written one. Staging the content in a sibling temporary file + and moving it over the destination keeps the previous document readable until the + new one is complete. The temporary file does not carry the document extension, so + a crash between the two steps cannot leave something that listing would pick up. + + Raises: + OSError: If the document could not be written. + """ + descriptor, temporary_name = tempfile.mkstemp(dir=str(path.parent), prefix=f".{path.name}.", suffix=".tmp") + temporary_path = Path(temporary_name) + try: + with os.fdopen(descriptor, "wb") as file: + file.write(content) + file.flush() + os.fsync(file.fileno()) + os.replace(temporary_path, path) + except BaseException: + temporary_path.unlink(missing_ok=True) + raise def _delete_document(self, name: str) -> None: """ @@ -195,14 +242,16 @@ def _local_directory(self, *, create: bool = False) -> Path: directory.mkdir(parents=True, exist_ok=True) return directory - def _list_blob_documents(self) -> dict[str, str]: + def _list_blob_documents(self) -> dict[str, bytes]: """ - Read documents from the configured Azure Blob container. + Read addressable documents from the configured Azure Blob container. Returns: - dict[str, str]: Document content keyed by blob stem. + dict[str, bytes]: Document content keyed by blob stem. """ - documents: dict[str, str] = {} + from azure.core.exceptions import AzureError + + documents: dict[str, bytes] = {} with self._open_container_client() as client: prefix = f"{self._blob_prefix}/" if self._blob_prefix else None blobs = client.list_blobs(name_starts_with=prefix) if prefix else client.list_blobs() @@ -210,8 +259,13 @@ def _list_blob_documents(self) -> dict[str, str]: blob.name for blob in blobs if self._is_direct_document_blob(blob_name=blob.name, prefix=prefix) ) for blob_name in blob_names: - relative_name = blob_name.removeprefix(prefix or "") - documents[PurePosixPath(relative_name).stem] = client.download_blob(blob_name).readall().decode("utf-8") + name = PurePosixPath(blob_name.removeprefix(prefix or "")).stem + if not self._is_addressable_name(name): + continue + try: + documents[name] = client.download_blob(blob_name).readall() + except AzureError as error: + logger.warning(f"Skipping unreadable document '{blob_name}' in {self.display_source}: {error}") return documents def _parse_blob_source(self) -> tuple[str, str]: diff --git a/pyrit/registry/scenario_preset_storage.py b/pyrit/registry/scenario_preset_storage.py index b52304303e..99352b6b63 100644 --- a/pyrit/registry/scenario_preset_storage.py +++ b/pyrit/registry/scenario_preset_storage.py @@ -36,14 +36,14 @@ def __init__(self, *, name: str, expected_version: str | None, actual_version: s super().__init__(f"Scenario preset '{name}' could not be saved because {detail}. Reload it and reapply.") -def _document_version(content: str) -> str: +def _document_version(content: bytes) -> str: """ Create an opaque version token from stored document content. Returns: str: The document-state version token. """ - return hashlib.sha256(content.encode()).hexdigest() + return hashlib.sha256(content).hexdigest() class ScenarioPresetStorage(FileDocumentStorage): @@ -114,8 +114,8 @@ def list_presets(self) -> dict[str, StoredPreset]: """ Read every stored preset, skipping any that cannot be parsed. - A malformed or hand-edited file must not prevent the rest of the library from - loading, so failures are logged and that preset is omitted. + A malformed, unreadable, or hand-edited file must not prevent the rest of the + library from loading, so failures are logged and that preset is omitted. Returns: dict[str, StoredPreset]: Stored presets keyed by name. @@ -138,7 +138,7 @@ def load_preset(self, name: str) -> StoredPreset | None: Raises: ValueError: If *name* is not a legal preset name. """ - content = self._read_document(name) + content = self._read_document_bytes(name) if content is None: return None preset = self._parse_preset(name=name, content=content) @@ -153,7 +153,8 @@ def get_preset_version(self, name: str) -> str | None: A document that cannot be parsed is otherwise unreachable: ``list_presets`` skips it, ``load_preset`` returns ``None``, and a create is refused because the document exists. Exposing its version lets a caller offer to overwrite the broken file - instead of leaving the name permanently unusable. + instead of leaving the name permanently unusable. The document is never decoded + here, so a file that is not even valid text can still be replaced. Returns: str | None: The document version, or ``None`` if no document is stored. @@ -161,7 +162,7 @@ def get_preset_version(self, name: str) -> str | None: Raises: ValueError: If *name* is not a legal preset name. """ - content = self._read_document(name) + content = self._read_document_bytes(name) return None if content is None else _document_version(content) def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) -> StoredPreset: @@ -187,7 +188,7 @@ def save_preset(self, *, preset: ScenarioPreset, expected_version: str | None) - ScenarioPresetConflictError: If the stored version does not match *expected_version*. ValueError: If the preset name is not a legal preset name. """ - existing_content = self._read_document(preset.name) + existing_content = self._read_document_bytes(preset.name) actual_version = None if existing_content is None else _document_version(existing_content) if actual_version != expected_version: raise ScenarioPresetConflictError( @@ -208,7 +209,7 @@ def delete_preset(self, name: str) -> None: self._delete_document(name) @staticmethod - def _serialize_preset(preset: ScenarioPreset) -> str: + def _serialize_preset(preset: ScenarioPreset) -> bytes: """ Serialize a preset to stored JSON. @@ -218,13 +219,13 @@ def _serialize_preset(preset: ScenarioPreset) -> str: expect a rename that cannot happen. Returns: - str: JSON document content. + bytes: Encoded JSON document content. """ payload = preset.model_dump(mode="json", exclude_none=True, exclude={"name"}) - return json.dumps(payload, indent=2, sort_keys=True) + "\n" + return (json.dumps(payload, indent=2, sort_keys=True) + "\n").encode("utf-8") @staticmethod - def _parse_preset(*, name: str, content: str) -> ScenarioPreset | None: + def _parse_preset(*, name: str, content: bytes) -> ScenarioPreset | None: """ Parse one stored preset document. @@ -236,7 +237,13 @@ def _parse_preset(*, name: str, content: str) -> ScenarioPreset | None: ScenarioPreset | None: The parsed preset, or ``None`` if it is malformed. """ try: - payload = json.loads(content) + text = content.decode("utf-8") + except UnicodeDecodeError: + logger.warning(f"Skipping stored scenario preset '{name}': it is not valid UTF-8 text.") + return None + + try: + payload = json.loads(text) except ValueError: logger.exception(f"Skipping stored scenario preset '{name}': it is not valid JSON.") return None diff --git a/tests/unit/registry/test_custom_initializer_storage.py b/tests/unit/registry/test_custom_initializer_storage.py index 3da0541888..bffed96216 100644 --- a/tests/unit/registry/test_custom_initializer_storage.py +++ b/tests/unit/registry/test_custom_initializer_storage.py @@ -148,6 +148,37 @@ def test_listing_skips_scripts_it_cannot_address(tmp_path: Path) -> None: storage.get_script_source(name) +def test_listing_skips_a_script_that_is_not_text(tmp_path: Path) -> None: + """Test that one undecodable file does not hide every other stored initializer.""" + (tmp_path / "good_one.py").write_bytes(b"VALUE = 1\n") + (tmp_path / "broken.py").write_bytes(b"\xff\xfe VALUE = 1") + storage = CustomInitializerStorage(source=str(tmp_path)) + + assert storage.list_scripts() == {"good_one": "VALUE = 1\n"} + + +def test_stored_crlf_script_is_normalized_before_it_is_handed_back(tmp_path: Path) -> None: + """Test that stored CRLF cannot round-trip into the CRCRLF a text-mode rewrite would produce.""" + (tmp_path / "windows_authored.py").write_bytes(b"VALUE = 1\r\nOTHER = 2\r\n") + storage = CustomInitializerStorage(source=str(tmp_path)) + + assert storage.list_scripts() == {"windows_authored": "VALUE = 1\nOTHER = 2\n"} + + +def test_blob_listing_skips_a_script_that_is_not_text() -> None: + """Test that one undecodable blob does not hide every other stored initializer.""" + client = MagicMock() + client.__enter__.return_value = client + client.list_blobs.return_value = [SimpleNamespace(name="broken.py"), SimpleNamespace(name="good_one.py")] + client.download_blob.side_effect = lambda blob_name: SimpleNamespace( + readall=lambda: b"\xff\xfe" if blob_name == "broken.py" else b"VALUE = 1\n" + ) + storage = CustomInitializerStorage(source="https://account.blob.core.windows.net/initializers?sig=secret") + + with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): + assert storage.list_scripts() == {"good_one": "VALUE = 1\n"} + + def test_direct_python_blob_rejects_name_outside_prefix(tmp_path: Path) -> None: """Test that blobs outside the configured virtual directory are ignored.""" storage = CustomInitializerStorage(source=str(tmp_path)) diff --git a/tests/unit/registry/test_scenario_preset_storage.py b/tests/unit/registry/test_scenario_preset_storage.py index 8bf4a140b5..9a0e70ec2c 100644 --- a/tests/unit/registry/test_scenario_preset_storage.py +++ b/tests/unit/registry/test_scenario_preset_storage.py @@ -3,7 +3,9 @@ """Tests for scenario preset storage.""" +import hashlib import json +import os from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -404,3 +406,121 @@ def test_listing_skips_documents_it_cannot_address(tmp_path: Path) -> None: (tmp_path / "My-Preset.json").write_text(json.dumps({"scenario_name": "foundry.red_team_agent"}), encoding="utf-8") assert sorted(storage.list_presets()) == ["nightly"] + + +def test_document_that_is_not_text_does_not_hide_valid_presets(tmp_path: Path) -> None: + """Test that a file which is not UTF-8 is skipped rather than failing the whole listing.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(), expected_version=None) + (tmp_path / "broken.json").write_bytes(b"\xff\xfe not utf-8") + + assert sorted(storage.list_presets()) == ["nightly"] + assert storage.load_preset("broken") is None + + +def test_unreadable_document_with_an_ignorable_name_is_never_read(tmp_path: Path) -> None: + """Test that a name the storage would refuse is skipped before its bytes are ever touched.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + storage.save_preset(preset=_make_preset(), expected_version=None) + (tmp_path / "My-Preset.json").write_bytes(b"\xff") + + assert sorted(storage.list_presets()) == ["nightly"] + + +def test_version_is_readable_for_a_document_that_is_not_text(tmp_path: Path) -> None: + """Test that the recovery path reaches a file too broken to decode, not just too broken to parse.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + (tmp_path / "nightly.json").write_bytes(b"\xff\xfe not utf-8") + + version = storage.get_preset_version("nightly") + assert version is not None + + storage.save_preset(preset=_make_preset(), expected_version=version) + + loaded = storage.load_preset("nightly") + assert loaded is not None + + +def test_version_matches_the_bytes_on_disk(tmp_path: Path) -> None: + """Test that the token describes stored bytes, so a later read cannot disagree with the save.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + + saved = storage.save_preset(preset=_make_preset(), expected_version=None) + + assert saved.version == hashlib.sha256((tmp_path / "nightly.json").read_bytes()).hexdigest() + assert storage.get_preset_version("nightly") == saved.version + + +def test_failed_write_preserves_the_previous_preset(tmp_path: Path) -> None: + """Test that a write that dies partway leaves the stored preset readable instead of empty.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + first = storage.save_preset(preset=_make_preset(description="first"), expected_version=None) + + def fail(source: object, target: object) -> None: + raise OSError(28, "No space left on device") + + with patch("os.replace", fail): + with pytest.raises(OSError): + storage.save_preset(preset=_make_preset(description="second"), expected_version=first.version) + + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.preset.description == "first" + assert loaded.version == first.version + assert sorted(path.name for path in tmp_path.iterdir()) == ["nightly.json"] + + +def test_a_reader_sees_the_previous_document_until_the_write_completes(tmp_path: Path) -> None: + """Test that an update never truncates the destination, so a reader sees the old or new document.""" + storage = ScenarioPresetStorage(source=str(tmp_path)) + first = storage.save_preset(preset=_make_preset(description="first"), expected_version=None) + observed: list[str] = [] + listed: list[list[str]] = [] + real_replace = os.replace + + def observe_then_replace(source: object, target: object) -> None: + observed.append((tmp_path / "nightly.json").read_text(encoding="utf-8")) + listed.append(sorted(path.name for path in tmp_path.glob("*.json"))) + real_replace(source, target) # type: ignore[arg-type] + + with patch("os.replace", observe_then_replace): + storage.save_preset(preset=_make_preset(description="second"), expected_version=first.version) + + assert json.loads(observed[0])["description"] == "first" + assert listed == [["nightly.json"]] + loaded = storage.load_preset("nightly") + assert loaded is not None + assert loaded.preset.description == "second" + + +def test_blob_listing_skips_a_document_that_is_not_text() -> None: + """Test that one undecodable blob does not hide every other stored preset.""" + document = json.dumps({"scenario_name": "foundry.red_team_agent"}).encode("utf-8") + client = MagicMock() + client.__enter__.return_value = client + client.list_blobs.return_value = [SimpleNamespace(name="broken.json"), SimpleNamespace(name="nightly.json")] + client.download_blob.side_effect = lambda blob_name: SimpleNamespace( + readall=lambda: b"\xff\xfe" if blob_name == "broken.json" else document + ) + storage = ScenarioPresetStorage(source="https://account.blob.core.windows.net/presets?sig=secret") + + with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): + presets = storage.list_presets() + + assert sorted(presets) == ["nightly"] + + +def test_blob_listing_does_not_download_a_name_it_would_refuse() -> None: + """Test that an ignorable blob name is filtered before the download that could fail on it.""" + document = json.dumps({"scenario_name": "foundry.red_team_agent"}).encode("utf-8") + client = MagicMock() + client.__enter__.return_value = client + client.list_blobs.return_value = [SimpleNamespace(name="My-Preset.json"), SimpleNamespace(name="nightly.json")] + client.download_blob.return_value.readall.return_value = document + storage = ScenarioPresetStorage(source="https://account.blob.core.windows.net/presets?sig=secret") + + with patch("azure.storage.blob.ContainerClient.from_container_url", return_value=client): + presets = storage.list_presets() + + assert sorted(presets) == ["nightly"] + assert [call.args[0] for call in client.download_blob.call_args_list] == ["nightly.json"] From cbe97dd17807679102618151e6a9534428f670ea Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Fri, 2 Oct 2026 14:07:01 -0400 Subject: [PATCH 06/16] FEAT: Add scenario preset CRUD and launch-resolution API Exposes the preset persistence layer over REST so a preset can be created, read, edited, deleted, and turned into a run request. Reads are open; mutations require an admin. Unlike custom initializers, presets are data rather than executable code, so they are not gated behind an allow_ switch. POST creates a preset that must not already exist and PUT requires the version returned when the preset was read, so a concurrent edit is reported as a 409 instead of being silently overwritten. References that do not resolve against the live registry are reported as advisory issues on the response rather than rejected, keeping a preset authored on one deployment storable on another. POST /{name}/resolve merges a preset with the launch-owned fields it omits and returns an ordinary RunScenarioRequest for the existing run endpoint. Resolving server-side keeps one implementation of the merge, so the distinction between unset and set-to-the-default cannot drift between clients, and leaves the launch path itself untouched. Storage is synchronous file or blob I/O, so every call into it runs on a worker thread rather than blocking the event loop. Adds scenario_presets_source so presets can be persisted to a shared Azure Blob container instead of the default local directory. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/getting_started/pyrit_conf.md | 10 + pyrit/backend/main.py | 2 + pyrit/backend/models/__init__.py | 12 + pyrit/backend/models/scenario_presets.py | 97 +++++ pyrit/backend/routes/__init__.py | 2 + pyrit/backend/routes/scenario_presets.py | 258 ++++++++++++++ pyrit/backend/services/__init__.py | 3 + pyrit/backend/services/runtime_lifecycle.py | 2 + .../services/scenario_preset_service.py | 313 +++++++++++++++++ pyrit/setup/configuration_loader.py | 33 +- .../backend/test_scenario_preset_routes.py | 288 +++++++++++++++ .../backend/test_scenario_preset_service.py | 330 ++++++++++++++++++ 12 files changed, 1346 insertions(+), 4 deletions(-) create mode 100644 pyrit/backend/models/scenario_presets.py create mode 100644 pyrit/backend/routes/scenario_presets.py create mode 100644 pyrit/backend/services/scenario_preset_service.py create mode 100644 tests/unit/backend/test_scenario_preset_routes.py create mode 100644 tests/unit/backend/test_scenario_preset_service.py diff --git a/doc/getting_started/pyrit_conf.md b/doc/getting_started/pyrit_conf.md index 20f922c9bd..da2d4b2e52 100644 --- a/doc/getting_started/pyrit_conf.md +++ b/doc/getting_started/pyrit_conf.md @@ -164,6 +164,16 @@ custom_initializers_source: https://account.blob.core.windows.net/pyrit-storage/ With this configuration, PyRIT reads and writes scripts directly under the `custom-initializers/` prefix in the `pyrit-storage` container. A SAS query string may be included; otherwise, PyRIT uses `DefaultAzureCredential`. The default is `~/.pyrit/custom_initializers`. +### `scenario_presets_source` + +Stores scenario presets in a local directory or Azure Blob container. An Azure URI may include a blob-name prefix, which behaves like a folder: + +```yaml +scenario_presets_source: https://account.blob.core.windows.net/pyrit-storage/scenario-presets +``` + +Point several deployments at the same container to share presets between them. A SAS query string may be included; otherwise, PyRIT uses `DefaultAzureCredential`. The default is `~/.pyrit/scenario_presets`. + ### `initialization_scripts` Local paths to custom Python scripts containing `PyRITInitializer` subclasses. Paths can be absolute or relative to the current working directory. diff --git a/pyrit/backend/main.py b/pyrit/backend/main.py index 0be8af50ca..9cf7c2086f 100644 --- a/pyrit/backend/main.py +++ b/pyrit/backend/main.py @@ -35,6 +35,7 @@ initializers, labels, media, + scenario_presets, scenarios, scores, targets, @@ -120,6 +121,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: app.include_router(converters.router, prefix="/api", tags=["converters"]) app.include_router(datasets.router, prefix="/api", tags=["datasets"]) app.include_router(scenarios.router, prefix="/api", tags=["scenarios"]) +app.include_router(scenario_presets.router, prefix="/api", tags=["scenario-presets"]) app.include_router(initializers.router, prefix="/api", tags=["initializers"]) app.include_router(labels.router, prefix="/api", tags=["labels"]) app.include_router(health.router, prefix="/api", tags=["health"]) diff --git a/pyrit/backend/models/__init__.py b/pyrit/backend/models/__init__.py index 827cead6f3..ad6c0cf706 100644 --- a/pyrit/backend/models/__init__.py +++ b/pyrit/backend/models/__init__.py @@ -62,6 +62,13 @@ ListRegisteredInitializersResponse, RegisterInitializerRequest, ) + from pyrit.backend.models.scenario_presets import ( + PresetIssue, + ResolveScenarioPresetRequest, + ScenarioPresetListResponse, + ScenarioPresetResponse, + UpdateScenarioPresetRequest, + ) from pyrit.backend.models.scenarios import ListRegisteredScenariosResponse, ScenarioRunListResponse from pyrit.backend.models.targets import CreateTargetRequest, TargetListResponse @@ -106,6 +113,11 @@ "PreviewStep": "pyrit.backend.models.converters", "DatasetInfo": "pyrit.backend.models.datasets", "DatasetListResponse": "pyrit.backend.models.datasets", + "PresetIssue": "pyrit.backend.models.scenario_presets", + "ResolveScenarioPresetRequest": "pyrit.backend.models.scenario_presets", + "ScenarioPresetListResponse": "pyrit.backend.models.scenario_presets", + "ScenarioPresetResponse": "pyrit.backend.models.scenario_presets", + "UpdateScenarioPresetRequest": "pyrit.backend.models.scenario_presets", "ListRegisteredScenariosResponse": "pyrit.backend.models.scenarios", "ScenarioRunListResponse": "pyrit.backend.models.scenarios", "ListRegisteredInitializersResponse": "pyrit.backend.models.initializers", diff --git a/pyrit/backend/models/scenario_presets.py b/pyrit/backend/models/scenario_presets.py new file mode 100644 index 0000000000..1aa1127b86 --- /dev/null +++ b/pyrit/backend/models/scenario_presets.py @@ -0,0 +1,97 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +REST envelopes for the scenario preset endpoints. + +The canonical preset types (``ScenarioPreset``, ``StoredPreset``) live in +``pyrit.models.catalog.scenario_preset`` and are imported from there directly. +The models here only add the wire-level concerns: the storage source, the +advisory issues attached to a read, and the launch-owned fields that a preset +deliberately does not carry. +""" + +from typing import Any + +from pydantic import BaseModel, Field + +from pyrit.models.catalog.scenario_preset import ScenarioPreset + +__all__ = [ + "PresetIssue", + "ResolveScenarioPresetRequest", + "ScenarioPresetListResponse", + "ScenarioPresetResponse", + "UpdateScenarioPresetRequest", +] + + +class PresetIssue(BaseModel): + """ + One advisory problem found while checking a preset against the live registry. + + Issues are reported, never enforced. A preset that names a scenario this + deployment has not registered is still a valid preset; it simply cannot run + here yet. Blocking the save would make presets un-portable between + deployments, which is the one property they exist to have. + """ + + field: str = Field(..., description="Preset field the issue applies to, e.g. 'scenario_name' or 'techniques'") + message: str = Field(..., description="Human-readable description of why the reference does not resolve") + + +class ScenarioPresetResponse(BaseModel): + """A stored preset, the version of the document it was read from, and its advisory issues.""" + + preset: ScenarioPreset = Field(..., description="The stored preset") + version: str = Field(..., description="Opaque version of the stored document; required to update it") + issues: list[PresetIssue] = Field( + default_factory=list, + description="Advisory problems resolving this preset against the live registry; empty when it is runnable", + ) + + +class ScenarioPresetListResponse(BaseModel): + """The configured preset storage source and every preset readable from it.""" + + source: str = Field(..., description="Credential-free configured preset source") + items: list[ScenarioPresetResponse] = Field(..., description="Stored presets, sorted by name") + + +class UpdateScenarioPresetRequest(BaseModel): + """ + Request body for updating an existing preset. + + ``expected_version`` is required rather than optional because updating is a + distinct operation from creating: a client that has not read the document it + is replacing has nothing to be optimistic about. Creation goes through POST, + where no version exists to supply. + """ + + preset: ScenarioPreset = Field(..., description="The replacement preset") + expected_version: str = Field(..., description="Version returned when the preset being edited was read") + + +class ResolveScenarioPresetRequest(BaseModel): + """ + The launch-owned fields a preset deliberately omits. + + A preset answers *what to test*; these answer *how and where*. Resolution is + a union of the two, which is why nothing here overlaps a preset field. + """ + + target_name: str = Field(..., description="Name of a registered target from the TargetRegistry") + adversarial_target_name: str | None = Field( + None, description="Name of a registered adversarial target, when the scenario uses one" + ) + initializers: list[str] | None = Field(None, description="Initializer names to run before the scenario") + initializer_args: dict[str, dict[str, Any]] | None = Field( + None, description="Per-initializer parameter overrides keyed by initializer name" + ) + max_concurrency: int | None = Field( + None, ge=1, le=100, description="Maximum concurrent operations; omit to use the run default" + ) + max_retries: int | None = Field( + None, ge=0, le=20, description="Maximum retry attempts on failure; omit to use the run default" + ) + labels: dict[str, str] | None = Field(None, description="Labels to attach to memory entries") diff --git a/pyrit/backend/routes/__init__.py b/pyrit/backend/routes/__init__.py index 454f2b4111..2b7df90c02 100644 --- a/pyrit/backend/routes/__init__.py +++ b/pyrit/backend/routes/__init__.py @@ -19,6 +19,7 @@ initializers, labels, media, + scenario_presets, scenarios, targets, version, @@ -32,6 +33,7 @@ "initializers": ("pyrit.backend.routes.initializers", None), "labels": ("pyrit.backend.routes.labels", None), "media": ("pyrit.backend.routes.media", None), + "scenario_presets": ("pyrit.backend.routes.scenario_presets", None), "scenarios": ("pyrit.backend.routes.scenarios", None), "targets": ("pyrit.backend.routes.targets", None), "version": ("pyrit.backend.routes.version", None), diff --git a/pyrit/backend/routes/scenario_presets.py b/pyrit/backend/routes/scenario_presets.py new file mode 100644 index 0000000000..0266cc1813 --- /dev/null +++ b/pyrit/backend/routes/scenario_presets.py @@ -0,0 +1,258 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Scenario preset API routes. + +Presets are data, not code, so unlike custom initializers they are not gated behind +an ``allow_`` switch. Reads are open and mutations require an admin, matching how +the rest of the API treats shared, server-side configuration. + +Route structure: + GET /api/scenario-presets — list every stored preset + GET /api/scenario-presets/{name} — get one preset and its version + POST /api/scenario-presets — create a preset that must not exist + PUT /api/scenario-presets/{name} — replace a preset the caller has read + DELETE /api/scenario-presets/{name} — delete a preset + POST /api/scenario-presets/{name}/resolve — combine a preset with launch fields +""" + +from azure.core.exceptions import AzureError +from fastapi import APIRouter, Depends, HTTPException, status + +from pyrit.backend.middleware.auth import require_admin +from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.scenario_presets import ( + ResolveScenarioPresetRequest, + ScenarioPresetListResponse, + ScenarioPresetResponse, + UpdateScenarioPresetRequest, +) +from pyrit.backend.services.scenario_preset_service import ( + ScenarioPresetNotFoundError, + get_scenario_preset_service, +) +from pyrit.models.catalog import RunScenarioRequest, ScenarioPreset +from pyrit.registry import ScenarioPresetConflictError + +router = APIRouter(prefix="/scenario-presets", tags=["scenario-presets"]) + + +def _storage_unavailable() -> HTTPException: + """ + Create a sanitized response for unavailable preset storage. + + Returns: + HTTPException: A service-unavailable response without SDK details. + """ + return HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Scenario preset storage is temporarily unavailable", + ) + + +def _not_found(name: str) -> HTTPException: + """ + Create a not-found response for a missing preset. + + Args: + name: The requested preset name. + + Returns: + HTTPException: A not-found response naming the preset. + """ + return HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Scenario preset '{name}' not found") + + +async def _load_preset_or_404_async(name: str) -> ScenarioPresetResponse: + """ + Read one preset, translating storage failures into HTTP responses. + + Args: + name: The preset name. + + Returns: + ScenarioPresetResponse: The stored preset, its version, and its advisory issues. + + Raises: + HTTPException: 404 if no preset is stored, 400 for an illegal name, 503 if storage fails. + """ + try: + preset = await get_scenario_preset_service().get_preset_async(name=name) + except AzureError as exc: + raise _storage_unavailable() from exc + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from None + + if preset is None: + raise _not_found(name) + return preset + + +@router.get("", response_model=ScenarioPresetListResponse) +async def list_scenario_presets() -> ScenarioPresetListResponse: # pyrit-async-suffix-exempt + """ + List every readable preset from the configured storage source. + + Presets that cannot be parsed are skipped rather than failing the listing. + + Returns: + ScenarioPresetListResponse: The configured source and the stored presets. + """ + try: + return await get_scenario_preset_service().list_presets_async() + except AzureError as exc: + raise _storage_unavailable() from exc + + +@router.get( + "/{name}", + response_model=ScenarioPresetResponse, + responses={404: {"model": ProblemDetail, "description": "Preset not found"}}, +) +async def get_scenario_preset(name: str) -> ScenarioPresetResponse: # pyrit-async-suffix-exempt + """ + Get one preset and the version required to update it. + + Args: + name: The preset name. + + Returns: + ScenarioPresetResponse: The stored preset, its version, and its advisory issues. + """ + return await _load_preset_or_404_async(name) + + +@router.post( + "", + response_model=ScenarioPresetResponse, + status_code=status.HTTP_201_CREATED, + dependencies=[Depends(require_admin)], + responses={ + 409: {"model": ProblemDetail, "description": "A preset is already stored under this name"}, + }, +) +async def create_scenario_preset(preset: ScenarioPreset) -> ScenarioPresetResponse: # pyrit-async-suffix-exempt + """ + Create a preset that must not already exist. + + References that do not resolve in this deployment are reported on the response + rather than rejected, so a preset authored elsewhere can still be stored here. + + Args: + preset: The preset to create. + + Returns: + ScenarioPresetResponse: The persisted preset, its version, and its advisory issues. + """ + try: + return await get_scenario_preset_service().save_preset_async(preset=preset, expected_version=None) + except ScenarioPresetConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Scenario preset '{preset.name}' already exists", + ) from exc + except AzureError as exc: + raise _storage_unavailable() from exc + + +@router.put( + "/{name}", + response_model=ScenarioPresetResponse, + dependencies=[Depends(require_admin)], + responses={ + 400: {"model": ProblemDetail, "description": "Body name does not match the path"}, + 409: {"model": ProblemDetail, "description": "The preset changed after it was read"}, + }, +) +async def update_scenario_preset( # pyrit-async-suffix-exempt + name: str, + body: UpdateScenarioPresetRequest, +) -> ScenarioPresetResponse: + """ + Replace a preset the caller has read. + + The update fails if the stored document no longer matches ``expected_version``, + so a concurrent edit is reported instead of being silently overwritten. + + Args: + name: The preset name from the path, which is authoritative. + body: The replacement preset and the version it is replacing. + + Returns: + ScenarioPresetResponse: The persisted preset, its new version, and its advisory issues. + + Raises: + HTTPException: 400 if the body names a different preset, 409 if the stored version moved. + """ + if body.preset.name != name: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Preset name '{body.preset.name}' does not match path '{name}'", + ) + + try: + return await get_scenario_preset_service().save_preset_async( + preset=body.preset, expected_version=body.expected_version + ) + except ScenarioPresetConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Scenario preset '{name}' changed since it was read; re-read it and retry", + ) from exc + except AzureError as exc: + raise _storage_unavailable() from exc + + +@router.delete( + "/{name}", + status_code=status.HTTP_204_NO_CONTENT, + dependencies=[Depends(require_admin)], + responses={404: {"model": ProblemDetail, "description": "Preset not found"}}, +) +async def delete_scenario_preset(name: str) -> None: # pyrit-async-suffix-exempt + """ + Delete one stored preset. + + Args: + name: The preset name. + + Raises: + HTTPException: 404 if no preset is stored, 400 for an illegal name, 503 if storage fails. + """ + try: + await get_scenario_preset_service().delete_preset_async(name=name) + except ScenarioPresetNotFoundError: + raise _not_found(name) from None + except AzureError as exc: + raise _storage_unavailable() from exc + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from None + + +@router.post( + "/{name}/resolve", + response_model=RunScenarioRequest, + responses={404: {"model": ProblemDetail, "description": "Preset not found"}}, +) +async def resolve_scenario_preset( # pyrit-async-suffix-exempt + name: str, + body: ResolveScenarioPresetRequest, +) -> RunScenarioRequest: + """ + Combine a stored preset with the launch-owned fields it omits. + + The result is an ordinary run request for the existing ``POST /scenarios/runs`` + endpoint. Resolving here rather than in each client keeps one implementation of + the merge, so the distinction between "unset" and "set to the default" cannot + drift between callers. + + Args: + name: The preset name. + body: The target and execution fields for this launch. + + Returns: + RunScenarioRequest: The request to post to the scenario run endpoint. + """ + stored = await _load_preset_or_404_async(name) + return get_scenario_preset_service().resolve_run_request(preset=stored.preset, launch=body) diff --git a/pyrit/backend/services/__init__.py b/pyrit/backend/services/__init__.py index 918bdfcdb0..ed5975bfcf 100644 --- a/pyrit/backend/services/__init__.py +++ b/pyrit/backend/services/__init__.py @@ -17,6 +17,7 @@ from pyrit.backend.services.converter_service import ConverterService, get_converter_service from pyrit.backend.services.dataset_service import DatasetService, get_dataset_service from pyrit.backend.services.initializer_service import InitializerService, get_initializer_service + from pyrit.backend.services.scenario_preset_service import ScenarioPresetService, get_scenario_preset_service from pyrit.backend.services.scenario_run_service import ScenarioRunService, get_scenario_run_service from pyrit.backend.services.scenario_service import ScenarioService, get_scenario_service from pyrit.backend.services.target_service import TargetService, get_target_service @@ -32,6 +33,8 @@ "get_initializer_service": "pyrit.backend.services.initializer_service", "ScenarioService": "pyrit.backend.services.scenario_service", "get_scenario_service": "pyrit.backend.services.scenario_service", + "ScenarioPresetService": "pyrit.backend.services.scenario_preset_service", + "get_scenario_preset_service": "pyrit.backend.services.scenario_preset_service", "ScenarioRunService": "pyrit.backend.services.scenario_run_service", "get_scenario_run_service": "pyrit.backend.services.scenario_run_service", "TargetService": "pyrit.backend.services.target_service", diff --git a/pyrit/backend/services/runtime_lifecycle.py b/pyrit/backend/services/runtime_lifecycle.py index 202ff2cb3f..04aea5e830 100644 --- a/pyrit/backend/services/runtime_lifecycle.py +++ b/pyrit/backend/services/runtime_lifecycle.py @@ -14,6 +14,7 @@ from pyrit.backend.models.initializers import ConfiguredInitializerSetting from pyrit.backend.services.configuration_file_service import ConfigurationFileService from pyrit.backend.services.environment_file_service import EnvironmentFileService +from pyrit.backend.services.scenario_preset_service import get_scenario_preset_service from pyrit.backend.services.scenario_run_service import get_scenario_run_service, peek_scenario_run_service from pyrit.backend.services.service_lifecycle import close_services_async, outstanding_estimates from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH @@ -88,6 +89,7 @@ async def _management_async(self, config: ConfigurationLoader) -> None: self.app.state.allow_custom_initializers = config.allow_custom_initializers registry = await asyncio.to_thread(InitializerRegistry.get_registry_singleton) registry.configure_custom_scripts_source(config.custom_initializers_source) + get_scenario_preset_service().configure_source(config.scenario_presets_source) def _publish(self, config: ConfigurationLoader) -> None: self.app.state.configured_initializers = [ diff --git a/pyrit/backend/services/scenario_preset_service.py b/pyrit/backend/services/scenario_preset_service.py new file mode 100644 index 0000000000..e1a7caa37c --- /dev/null +++ b/pyrit/backend/services/scenario_preset_service.py @@ -0,0 +1,313 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Scenario preset service for CRUD and launch resolution. +""" + +import asyncio +import logging +from functools import lru_cache +from typing import Any + +from pyrit.backend.models.scenario_presets import ( + PresetIssue, + ResolveScenarioPresetRequest, + ScenarioPresetListResponse, + ScenarioPresetResponse, +) +from pyrit.backend.services.scenario_service import get_scenario_service +from pyrit.models.catalog import RegisteredScenario, RunScenarioRequest, ScenarioPreset, StoredPreset +from pyrit.registry import ScenarioPresetStorage + +logger = logging.getLogger(__name__) + + +class ScenarioPresetNotFoundError(KeyError): + """No preset is stored under the requested name.""" + + +class ScenarioPresetService: + """ + Service for reading, writing, and resolving scenario presets. + + Storage is synchronous file or blob I/O, so every call into it is handed to a + worker thread. Running it inline would block the event loop for the whole + directory listing on every request. + """ + + def __init__(self) -> None: + """Initialize the service without touching storage.""" + self._storage: ScenarioPresetStorage | None = None + + def configure_source(self, source: str | None) -> None: + """ + Configure the local directory or Azure Blob source for scenario presets. + + Args: + source (str | None): The configured source, or None to use the default directory. + """ + self._storage = ScenarioPresetStorage(source=source) + + async def list_presets_async(self) -> ScenarioPresetListResponse: + """ + Read every stored preset, newest advisory issues included. + + Returns: + ScenarioPresetListResponse: The configured source and every readable preset. + """ + storage = self._get_storage() + stored = await asyncio.to_thread(storage.list_presets) + + items = [ + await self._to_response_async(stored_preset) for _, stored_preset in sorted(stored.items(), key=_by_name) + ] + return ScenarioPresetListResponse(source=storage.display_source, items=items) + + async def get_preset_async(self, *, name: str) -> ScenarioPresetResponse | None: + """ + Read one stored preset. + + Args: + name (str): The preset name. + + Returns: + ScenarioPresetResponse | None: The preset, or None if nothing is stored under *name*. + """ + stored = await asyncio.to_thread(self._get_storage().load_preset, name) + return None if stored is None else await self._to_response_async(stored) + + async def save_preset_async( + self, + *, + preset: ScenarioPreset, + expected_version: str | None, + ) -> ScenarioPresetResponse: + """ + Persist one preset. + + Unresolvable references are reported on the response rather than rejected, so a + preset authored against one deployment can be saved on another. + + Args: + preset (ScenarioPreset): The preset to persist. + expected_version (str | None): None to create a preset that must not already exist, + or the version returned when the edited preset was read. + + Returns: + ScenarioPresetResponse: The persisted preset, its new version, and its advisory issues. + + Raises: + ScenarioPresetConflictError: If the stored version does not match *expected_version*. + """ + stored = await asyncio.to_thread( + self._get_storage().save_preset, preset=preset, expected_version=expected_version + ) + logger.info("Saved scenario preset: %s", preset.name) + return await self._to_response_async(stored) + + async def delete_preset_async(self, *, name: str) -> None: + """ + Delete one stored preset. + + Args: + name (str): The preset name. + + Raises: + ScenarioPresetNotFoundError: If nothing is stored under *name*. + """ + storage = self._get_storage() + if await asyncio.to_thread(storage.get_preset_version, name) is None: + raise ScenarioPresetNotFoundError(name) + + await asyncio.to_thread(storage.delete_preset, name) + logger.info("Deleted scenario preset: %s", name) + + @staticmethod + def resolve_run_request(*, preset: ScenarioPreset, launch: ResolveScenarioPresetRequest) -> RunScenarioRequest: + """ + Combine a preset with the launch-owned fields it omits. + + ``max_concurrency`` and ``max_retries`` are the only run fields that are not + tri-state, so an unset value is dropped rather than passed as None. Passing it + through would pin the run default at resolution time instead of letting the + request model supply it. + + Args: + preset (ScenarioPreset): The preset supplying the scenario-owned fields. + launch (ResolveScenarioPresetRequest): The target and execution fields for this launch. + + Returns: + RunScenarioRequest: The request to post to the existing scenario run endpoint. + """ + run_defaults: dict[str, Any] = { + "max_concurrency": launch.max_concurrency, + "max_retries": launch.max_retries, + } + + return RunScenarioRequest( + scenario_name=preset.scenario_name, + techniques=preset.techniques, + dataset_names=preset.dataset_names, + max_dataset_size=preset.max_dataset_size, + dataset_filters=preset.dataset_filters, + include_baseline=preset.include_baseline, + scenario_params=preset.scenario_params, + target_name=launch.target_name, + adversarial_target_name=launch.adversarial_target_name, + initializers=launch.initializers, + initializer_args=launch.initializer_args, + labels=launch.labels, + **{name: value for name, value in run_defaults.items() if value is not None}, + ) + + def _get_storage(self) -> ScenarioPresetStorage: + """ + Return storage for the configured preset source. + + Returns: + ScenarioPresetStorage: The configured storage, defaulting to the standard directory. + """ + storage = self._storage + if storage is None: + storage = ScenarioPresetStorage() + self._storage = storage + return storage + + async def _to_response_async(self, stored: StoredPreset) -> ScenarioPresetResponse: + """ + Attach advisory issues to a stored preset. + + Args: + stored (StoredPreset): The preset and the version it was read from. + + Returns: + ScenarioPresetResponse: The wire representation of the preset. + """ + issues = await self._collect_issues_async(preset=stored.preset) + return ScenarioPresetResponse(preset=stored.preset, version=stored.version, issues=issues) + + async def _collect_issues_async(self, *, preset: ScenarioPreset) -> list[PresetIssue]: + """ + Check a preset against the live scenario registry. + + Args: + preset (ScenarioPreset): The preset to check. + + Returns: + list[PresetIssue]: Advisory issues, empty when the preset resolves here. + """ + scenario = await get_scenario_service().get_scenario_async(scenario_name=preset.scenario_name) + if scenario is None: + return [ + PresetIssue( + field="scenario_name", + message=f"Scenario '{preset.scenario_name}' is not registered in this deployment.", + ) + ] + + return [ + *_unknown_technique_issues(preset=preset, scenario=scenario), + *_unknown_parameter_issues(preset=preset, scenario=scenario), + *_forbidden_baseline_issues(preset=preset, scenario=scenario), + ] + + +def _by_name(item: tuple[str, StoredPreset]) -> str: + """ + Sort key for stored presets. + + Args: + item (tuple[str, StoredPreset]): A storage name and the preset read from it. + + Returns: + str: The storage name. + """ + return item[0] + + +def _unknown_technique_issues(*, preset: ScenarioPreset, scenario: RegisteredScenario) -> list[PresetIssue]: + """ + Report techniques the scenario does not expose. + + Args: + preset (ScenarioPreset): The preset to check. + scenario (RegisteredScenario): The registered scenario it names. + + Returns: + list[PresetIssue]: One issue naming every unknown technique, or an empty list. + """ + if not preset.techniques: + return [] + + known = set(scenario.all_techniques) | set(scenario.aggregate_techniques) + unknown = [technique for technique in preset.techniques if technique not in known] + if not unknown: + return [] + + return [ + PresetIssue( + field="techniques", + message=f"Scenario '{scenario.scenario_name}' does not define: {', '.join(sorted(unknown))}.", + ) + ] + + +def _unknown_parameter_issues(*, preset: ScenarioPreset, scenario: RegisteredScenario) -> list[PresetIssue]: + """ + Report scenario parameters the scenario does not declare. + + Args: + preset (ScenarioPreset): The preset to check. + scenario (RegisteredScenario): The registered scenario it names. + + Returns: + list[PresetIssue]: One issue naming every undeclared parameter, or an empty list. + """ + if not preset.scenario_params: + return [] + + declared = {parameter.name for parameter in scenario.supported_parameters} + unknown = [name for name in preset.scenario_params if name not in declared] + if not unknown: + return [] + + return [ + PresetIssue( + field="scenario_params", + message=f"Scenario '{scenario.scenario_name}' does not declare: {', '.join(sorted(unknown))}.", + ) + ] + + +def _forbidden_baseline_issues(*, preset: ScenarioPreset, scenario: RegisteredScenario) -> list[PresetIssue]: + """ + Report a baseline request the scenario forbids. + + Args: + preset (ScenarioPreset): The preset to check. + scenario (RegisteredScenario): The registered scenario it names. + + Returns: + list[PresetIssue]: A single issue when the scenario forbids a requested baseline. + """ + if not preset.include_baseline or scenario.baseline_policy != "forbidden": + return [] + + return [ + PresetIssue( + field="include_baseline", + message=f"Scenario '{scenario.scenario_name}' does not support a baseline run.", + ) + ] + + +@lru_cache(maxsize=1) +def get_scenario_preset_service() -> ScenarioPresetService: + """ + Get the global scenario preset service instance. + + Returns: + ScenarioPresetService: The singleton scenario preset service instance. + """ + return ScenarioPresetService() diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index d821144c64..ca1c68a033 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -111,6 +111,8 @@ class ConfigurationLoader(YamlLoadable): bootstrap document should fail initialization. custom_initializers_source: Local directory or Azure Blob container URI, optionally followed by a blob prefix, used to persist custom initializer Python scripts. + scenario_presets_source: Local directory or Azure Blob container URI, optionally + followed by a blob prefix, used to persist scenario presets. silent: Whether to suppress initialization messages. seed: Optional root seed for deterministic converter operations. operator: Name for the current operator, e.g. a team or username. @@ -162,6 +164,7 @@ class ConfigurationLoader(YamlLoadable): enable_live_reinitialization: bool = False allow_custom_initializers: bool = False custom_initializers_source: str | None = None + scenario_presets_source: str | None = None server: dict[str, Any] | None = None extensions: dict[str, Any] = field(default_factory=dict) @@ -188,6 +191,7 @@ def __post_init__(self) -> None: self._normalize_initializers() self._validate_env_akv_ref() self._validate_custom_initializers_source() + self._validate_scenario_presets_source() self._normalize_server() def _validate_allow_custom_initializers(self) -> None: @@ -207,10 +211,31 @@ def _validate_custom_initializers_source(self) -> None: Raises: ValueError: If the source is not a non-empty string. """ - if self.custom_initializers_source is not None and ( - not isinstance(self.custom_initializers_source, str) or not self.custom_initializers_source.strip() - ): - raise ValueError("custom_initializers_source must be a non-empty local directory or container URI.") + self._validate_document_source(name="custom_initializers_source", value=self.custom_initializers_source) + + def _validate_scenario_presets_source(self) -> None: + """ + Validate the optional scenario preset storage source. + + Raises: + ValueError: If the source is not a non-empty string. + """ + self._validate_document_source(name="scenario_presets_source", value=self.scenario_presets_source) + + @staticmethod + def _validate_document_source(*, name: str, value: Any) -> None: + """ + Validate an optional document storage source. + + Args: + name: The configuration key being validated, used in the error message. + value: The configured value. + + Raises: + ValueError: If the source is not a non-empty string. + """ + if value is not None and (not isinstance(value, str) or not value.strip()): + raise ValueError(f"{name} must be a non-empty local directory or container URI.") def _validate_env_akv_ref(self) -> None: """ diff --git a/tests/unit/backend/test_scenario_preset_routes.py b/tests/unit/backend/test_scenario_preset_routes.py new file mode 100644 index 0000000000..4196b29003 --- /dev/null +++ b/tests/unit/backend/test_scenario_preset_routes.py @@ -0,0 +1,288 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the scenario preset routes.""" + +from collections.abc import Iterator +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from azure.core.exceptions import AzureError +from fastapi.testclient import TestClient + +from pyrit.backend.main import app +from pyrit.backend.middleware.auth import require_admin +from pyrit.backend.models.scenario_presets import ( + PresetIssue, + ScenarioPresetListResponse, + ScenarioPresetResponse, +) +from pyrit.backend.services.scenario_preset_service import ( + ScenarioPresetNotFoundError, + ScenarioPresetService, +) +from pyrit.models.catalog import ScenarioPreset +from pyrit.registry import ScenarioPresetConflictError + +PRESET_NAME = "quick_scan" +SCENARIO_NAME = "foundry.red_team_agent" + + +@pytest.fixture +def client(compatibility_headers: dict[str, str]) -> Iterator[TestClient]: + """Create a test client with admin authorization satisfied.""" + app.dependency_overrides[require_admin] = lambda: None + try: + yield TestClient(app, headers=compatibility_headers) + finally: + app.dependency_overrides.pop(require_admin, None) + + +@pytest.fixture +def anonymous_client(compatibility_headers: dict[str, str]) -> Iterator[TestClient]: + """Create a test client without admin authorization.""" + yield TestClient(app, headers=compatibility_headers) + + +@pytest.fixture +def service() -> Iterator[MagicMock]: + """Patch the preset service the routes resolve.""" + mock_service = MagicMock(spec=ScenarioPresetService) + with patch( + "pyrit.backend.routes.scenario_presets.get_scenario_preset_service", + return_value=mock_service, + ): + yield mock_service + + +def _preset(name: str = PRESET_NAME) -> ScenarioPreset: + """Build a minimal preset.""" + return ScenarioPreset(name=name, scenario_name=SCENARIO_NAME) + + +def _response(*, name: str = PRESET_NAME, version: str = "v1", issues: list[PresetIssue] | None = None): + """Build a preset response envelope.""" + return ScenarioPresetResponse(preset=_preset(name), version=version, issues=issues or []) + + +class TestListPresets: + """GET /api/scenario-presets.""" + + def test_list_returns_the_source_and_items(self, client: TestClient, service: MagicMock) -> None: + service.list_presets_async = AsyncMock( + return_value=ScenarioPresetListResponse(source="/tmp/presets", items=[_response()]) + ) + + response = client.get("/api/scenario-presets") + + assert response.status_code == 200 + body = response.json() + assert body["source"] == "/tmp/presets" + assert body["items"][0]["preset"]["name"] == PRESET_NAME + assert body["items"][0]["version"] == "v1" + + def test_list_reports_advisory_issues(self, client: TestClient, service: MagicMock) -> None: + issue = PresetIssue(field="scenario_name", message="Scenario is not registered in this deployment.") + service.list_presets_async = AsyncMock( + return_value=ScenarioPresetListResponse(source="/tmp", items=[_response(issues=[issue])]) + ) + + response = client.get("/api/scenario-presets") + + assert response.json()["items"][0]["issues"] == [ + {"field": "scenario_name", "message": "Scenario is not registered in this deployment."} + ] + + def test_list_is_readable_without_admin(self, anonymous_client: TestClient, service: MagicMock) -> None: + service.list_presets_async = AsyncMock(return_value=ScenarioPresetListResponse(source="/tmp", items=[])) + + assert anonymous_client.get("/api/scenario-presets").status_code == 200 + + def test_storage_failure_is_reported_without_sdk_detail(self, client: TestClient, service: MagicMock) -> None: + service.list_presets_async = AsyncMock(side_effect=AzureError("container 'x' key=secret")) + + response = client.get("/api/scenario-presets") + + assert response.status_code == 503 + assert "secret" not in response.text + + +class TestGetPreset: + """GET /api/scenario-presets/{name}.""" + + def test_get_returns_the_preset_and_version(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(return_value=_response()) + + response = client.get(f"/api/scenario-presets/{PRESET_NAME}") + + assert response.status_code == 200 + assert response.json()["version"] == "v1" + + def test_get_reports_a_missing_preset(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(return_value=None) + + response = client.get("/api/scenario-presets/missing") + + assert response.status_code == 404 + + def test_an_illegal_name_is_a_client_error(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(side_effect=ValueError("Invalid registry name 'Bad Name'")) + + response = client.get("/api/scenario-presets/BadName") + + assert response.status_code == 400 + + +class TestCreatePreset: + """POST /api/scenario-presets.""" + + def test_create_returns_201_with_the_new_version(self, client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock(return_value=_response(version="v2")) + + response = client.post("/api/scenario-presets", json=_preset().model_dump()) + + assert response.status_code == 201 + assert response.json()["version"] == "v2" + assert service.save_preset_async.call_args.kwargs["expected_version"] is None + + def test_create_over_an_existing_name_is_a_conflict(self, client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock( + side_effect=ScenarioPresetConflictError(name=PRESET_NAME, expected_version=None, actual_version="v1") + ) + + response = client.post("/api/scenario-presets", json=_preset().model_dump()) + + assert response.status_code == 409 + + def test_create_requires_admin(self, anonymous_client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock(return_value=_response()) + + response = anonymous_client.post("/api/scenario-presets", json=_preset().model_dump()) + + assert response.status_code == 403 + service.save_preset_async.assert_not_called() + + def test_an_unknown_field_is_rejected(self, client: TestClient, service: MagicMock) -> None: + payload = _preset().model_dump() + payload["tecniques"] = ["crescendo"] + + response = client.post("/api/scenario-presets", json=payload) + + assert response.status_code == 422 + + +class TestUpdatePreset: + """PUT /api/scenario-presets/{name}.""" + + def test_update_passes_the_expected_version_through(self, client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock(return_value=_response(version="v2")) + + response = client.put( + f"/api/scenario-presets/{PRESET_NAME}", + json={"preset": _preset().model_dump(), "expected_version": "v1"}, + ) + + assert response.status_code == 200 + assert service.save_preset_async.call_args.kwargs["expected_version"] == "v1" + + def test_a_body_naming_a_different_preset_is_rejected(self, client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock(return_value=_response()) + + response = client.put( + "/api/scenario-presets/other_name", + json={"preset": _preset().model_dump(), "expected_version": "v1"}, + ) + + assert response.status_code == 400 + service.save_preset_async.assert_not_called() + + def test_a_stale_version_is_a_conflict(self, client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock( + side_effect=ScenarioPresetConflictError(name=PRESET_NAME, expected_version="v1", actual_version="v2") + ) + + response = client.put( + f"/api/scenario-presets/{PRESET_NAME}", + json={"preset": _preset().model_dump(), "expected_version": "v1"}, + ) + + assert response.status_code == 409 + + def test_an_omitted_version_is_rejected(self, client: TestClient, service: MagicMock) -> None: + response = client.put( + f"/api/scenario-presets/{PRESET_NAME}", + json={"preset": _preset().model_dump()}, + ) + + assert response.status_code == 422 + + def test_update_requires_admin(self, anonymous_client: TestClient, service: MagicMock) -> None: + service.save_preset_async = AsyncMock(return_value=_response()) + + response = anonymous_client.put( + f"/api/scenario-presets/{PRESET_NAME}", + json={"preset": _preset().model_dump(), "expected_version": "v1"}, + ) + + assert response.status_code == 403 + service.save_preset_async.assert_not_called() + + +class TestDeletePreset: + """DELETE /api/scenario-presets/{name}.""" + + def test_delete_returns_204(self, client: TestClient, service: MagicMock) -> None: + service.delete_preset_async = AsyncMock(return_value=None) + + response = client.delete(f"/api/scenario-presets/{PRESET_NAME}") + + assert response.status_code == 204 + + def test_delete_reports_a_missing_preset(self, client: TestClient, service: MagicMock) -> None: + service.delete_preset_async = AsyncMock(side_effect=ScenarioPresetNotFoundError(PRESET_NAME)) + + response = client.delete(f"/api/scenario-presets/{PRESET_NAME}") + + assert response.status_code == 404 + + def test_delete_requires_admin(self, anonymous_client: TestClient, service: MagicMock) -> None: + service.delete_preset_async = AsyncMock(return_value=None) + + response = anonymous_client.delete(f"/api/scenario-presets/{PRESET_NAME}") + + assert response.status_code == 403 + service.delete_preset_async.assert_not_called() + + +class TestResolvePreset: + """POST /api/scenario-presets/{name}/resolve.""" + + def test_resolve_returns_a_run_request(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(return_value=_response()) + service.resolve_run_request = ScenarioPresetService.resolve_run_request + + response = client.post( + f"/api/scenario-presets/{PRESET_NAME}/resolve", + json={"target_name": "gpt4"}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["scenario_name"] == SCENARIO_NAME + assert body["target_name"] == "gpt4" + assert body["max_concurrency"] == 10 + assert body["techniques"] is None + + def test_resolve_reports_a_missing_preset(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(return_value=None) + + response = client.post("/api/scenario-presets/missing/resolve", json={"target_name": "gpt4"}) + + assert response.status_code == 404 + + def test_resolve_requires_a_target(self, client: TestClient, service: MagicMock) -> None: + service.get_preset_async = AsyncMock(return_value=_response()) + + response = client.post(f"/api/scenario-presets/{PRESET_NAME}/resolve", json={}) + + assert response.status_code == 422 diff --git a/tests/unit/backend/test_scenario_preset_service.py b/tests/unit/backend/test_scenario_preset_service.py new file mode 100644 index 0000000000..a20ead70f1 --- /dev/null +++ b/tests/unit/backend/test_scenario_preset_service.py @@ -0,0 +1,330 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the scenario preset service.""" + +from pathlib import Path +from unittest.mock import AsyncMock, patch + +import pytest + +from pyrit.backend.models.scenario_presets import ResolveScenarioPresetRequest +from pyrit.backend.services.scenario_preset_service import ( + ScenarioPresetNotFoundError, + ScenarioPresetService, + get_scenario_preset_service, +) +from pyrit.models import Parameter +from pyrit.models.catalog import RegisteredScenario, ScenarioPreset +from pyrit.models.catalog.scenario import ScenarioRunSizeEstimate +from pyrit.registry import ScenarioPresetConflictError + +SCENARIO_NAME = "foundry.red_team_agent" + + +def _registered_scenario( + *, + all_techniques: list[str] | None = None, + aggregate_techniques: list[str] | None = None, + supported_parameters: list[str] | None = None, + baseline_policy: str = "enabled", +) -> RegisteredScenario: + """Build a registered scenario with only the fields preset validation reads.""" + techniques = all_techniques if all_techniques is not None else ["crescendo", "flip"] + return RegisteredScenario( + scenario_name=SCENARIO_NAME, + scenario_type="RedTeamAgentScenario", + description="Test scenario", + default_technique=techniques[0], + all_techniques=techniques, + aggregate_techniques=aggregate_techniques or [], + default_datasets=["harmbench"], + baseline_policy=baseline_policy, + supported_parameters=[ + Parameter(name=name, description="", param_type=str) for name in (supported_parameters or ["max_turns"]) + ], + default_run_size=ScenarioRunSizeEstimate.unavailable(), + ) + + +@pytest.fixture +def service(tmp_path: Path) -> ScenarioPresetService: + """Create a service backed by an isolated preset directory.""" + instance = ScenarioPresetService() + instance.configure_source(str(tmp_path)) + return instance + + +@pytest.fixture +def registered_scenario() -> RegisteredScenario: + """Patch the scenario lookup so presets resolve against a known scenario.""" + scenario = _registered_scenario() + with patch("pyrit.backend.services.scenario_preset_service.get_scenario_service") as mock_factory: + mock_factory.return_value.get_scenario_async = AsyncMock(return_value=scenario) + yield scenario + + +def _preset(name: str = "quick_scan", **overrides: object) -> ScenarioPreset: + """Build a preset that resolves cleanly against the fixture scenario.""" + fields: dict[str, object] = {"name": name, "scenario_name": SCENARIO_NAME} + fields.update(overrides) + return ScenarioPreset(**fields) # type: ignore[arg-type] + + +class TestPresetCrud: + """CRUD behavior over the configured storage source.""" + + async def test_list_presets_is_empty_when_nothing_is_stored( + self, service: ScenarioPresetService, tmp_path: Path, registered_scenario: RegisteredScenario + ) -> None: + response = await service.list_presets_async() + + assert response.items == [] + assert response.source == str(tmp_path) + + async def test_list_presets_is_sorted_by_name( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + for name in ("zebra", "alpha", "middle"): + await service.save_preset_async(preset=_preset(name), expected_version=None) + + response = await service.list_presets_async() + + assert [item.preset.name for item in response.items] == ["alpha", "middle", "zebra"] + + async def test_get_preset_returns_none_when_not_stored( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + assert await service.get_preset_async(name="missing") is None + + async def test_saved_preset_round_trips_every_field( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + preset = _preset( + description="Nightly smoke", + techniques=["crescendo"], + dataset_names=["harmbench"], + max_dataset_size=25, + dataset_filters={"harm_categories": ["violence"]}, + include_baseline=True, + scenario_params={"max_turns": 3}, + ) + + await service.save_preset_async(preset=preset, expected_version=None) + read_back = await service.get_preset_async(name=preset.name) + + assert read_back is not None + assert read_back.preset == preset + + async def test_update_with_the_version_from_a_read_succeeds( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + created = await service.save_preset_async(preset=_preset(), expected_version=None) + + updated = await service.save_preset_async( + preset=_preset(description="changed"), expected_version=created.version + ) + + assert updated.preset.description == "changed" + assert updated.version != created.version + + async def test_update_with_a_stale_version_conflicts( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + created = await service.save_preset_async(preset=_preset(), expected_version=None) + await service.save_preset_async(preset=_preset(description="first"), expected_version=created.version) + + with pytest.raises(ScenarioPresetConflictError): + await service.save_preset_async(preset=_preset(description="second"), expected_version=created.version) + + async def test_create_over_an_existing_preset_conflicts( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + await service.save_preset_async(preset=_preset(), expected_version=None) + + with pytest.raises(ScenarioPresetConflictError): + await service.save_preset_async(preset=_preset(), expected_version=None) + + async def test_delete_removes_the_stored_preset( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + await service.save_preset_async(preset=_preset(), expected_version=None) + + await service.delete_preset_async(name="quick_scan") + + assert await service.get_preset_async(name="quick_scan") is None + + async def test_delete_reports_a_missing_preset_rather_than_succeeding_silently( + self, service: ScenarioPresetService + ) -> None: + with pytest.raises(ScenarioPresetNotFoundError): + await service.delete_preset_async(name="missing") + + +class TestAdvisoryValidation: + """Unresolvable references are reported, never enforced.""" + + async def test_a_resolvable_preset_has_no_issues( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + saved = await service.save_preset_async( + preset=_preset(techniques=["crescendo"], scenario_params={"max_turns": 2}), expected_version=None + ) + + assert saved.issues == [] + + async def test_an_unregistered_scenario_is_saved_and_reported(self, service: ScenarioPresetService) -> None: + with patch("pyrit.backend.services.scenario_preset_service.get_scenario_service") as mock_factory: + mock_factory.return_value.get_scenario_async = AsyncMock(return_value=None) + saved = await service.save_preset_async(preset=_preset(), expected_version=None) + + read_back = await service.get_preset_async(name="quick_scan") + + assert read_back is not None + assert [issue.field for issue in saved.issues] == ["scenario_name"] + assert "not registered" in saved.issues[0].message + + async def test_unknown_techniques_are_reported_without_blocking_the_save( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + saved = await service.save_preset_async( + preset=_preset(techniques=["crescendo", "not_a_technique"]), expected_version=None + ) + + assert [issue.field for issue in saved.issues] == ["techniques"] + assert "not_a_technique" in saved.issues[0].message + assert await service.get_preset_async(name="quick_scan") is not None + + async def test_an_aggregate_technique_is_not_reported_as_unknown(self, service: ScenarioPresetService) -> None: + scenario = _registered_scenario(aggregate_techniques=["all"]) + with patch("pyrit.backend.services.scenario_preset_service.get_scenario_service") as mock_factory: + mock_factory.return_value.get_scenario_async = AsyncMock(return_value=scenario) + saved = await service.save_preset_async(preset=_preset(techniques=["all"]), expected_version=None) + + assert saved.issues == [] + + async def test_undeclared_scenario_parameters_are_reported( + self, service: ScenarioPresetService, registered_scenario: RegisteredScenario + ) -> None: + saved = await service.save_preset_async( + preset=_preset(scenario_params={"max_turns": 1, "mystery": 2}), expected_version=None + ) + + assert [issue.field for issue in saved.issues] == ["scenario_params"] + assert "mystery" in saved.issues[0].message + + async def test_a_baseline_the_scenario_forbids_is_reported(self, service: ScenarioPresetService) -> None: + scenario = _registered_scenario(baseline_policy="forbidden") + with patch("pyrit.backend.services.scenario_preset_service.get_scenario_service") as mock_factory: + mock_factory.return_value.get_scenario_async = AsyncMock(return_value=scenario) + saved = await service.save_preset_async(preset=_preset(include_baseline=True), expected_version=None) + + assert [issue.field for issue in saved.issues] == ["include_baseline"] + + async def test_an_omitted_baseline_is_not_reported_when_the_scenario_forbids_one( + self, service: ScenarioPresetService + ) -> None: + scenario = _registered_scenario(baseline_policy="forbidden") + with patch("pyrit.backend.services.scenario_preset_service.get_scenario_service") as mock_factory: + mock_factory.return_value.get_scenario_async = AsyncMock(return_value=scenario) + saved = await service.save_preset_async(preset=_preset(), expected_version=None) + + assert saved.issues == [] + + +class TestRunRequestResolution: + """A preset plus launch fields becomes an ordinary run request.""" + + def test_unset_run_fields_fall_back_to_the_request_defaults(self) -> None: + resolved = ScenarioPresetService.resolve_run_request( + preset=_preset(), launch=ResolveScenarioPresetRequest(target_name="gpt4") + ) + + assert resolved.max_concurrency == 10 + assert resolved.max_retries == 0 + + def test_explicit_run_fields_are_applied(self) -> None: + resolved = ScenarioPresetService.resolve_run_request( + preset=_preset(), + launch=ResolveScenarioPresetRequest(target_name="gpt4", max_concurrency=4, max_retries=2), + ) + + assert resolved.max_concurrency == 4 + assert resolved.max_retries == 2 + + def test_unset_preset_fields_stay_unset_so_the_scenario_default_still_applies(self) -> None: + resolved = ScenarioPresetService.resolve_run_request( + preset=_preset(), launch=ResolveScenarioPresetRequest(target_name="gpt4") + ) + + assert resolved.techniques is None + assert resolved.dataset_names is None + assert resolved.max_dataset_size is None + assert resolved.dataset_filters is None + assert resolved.include_baseline is None + assert resolved.scenario_params is None + + def test_a_baseline_explicitly_disabled_by_the_preset_is_preserved(self) -> None: + resolved = ScenarioPresetService.resolve_run_request( + preset=_preset(include_baseline=False), launch=ResolveScenarioPresetRequest(target_name="gpt4") + ) + + assert resolved.include_baseline is False + + def test_every_preset_field_reaches_the_run_request(self) -> None: + preset = _preset( + techniques=["crescendo"], + dataset_names=["harmbench"], + max_dataset_size=25, + dataset_filters={"harm_categories": ["violence"]}, + include_baseline=True, + scenario_params={"max_turns": 3}, + ) + + resolved = ScenarioPresetService.resolve_run_request( + preset=preset, launch=ResolveScenarioPresetRequest(target_name="gpt4") + ) + + assert resolved.scenario_name == preset.scenario_name + assert resolved.techniques == preset.techniques + assert resolved.dataset_names == preset.dataset_names + assert resolved.max_dataset_size == preset.max_dataset_size + assert resolved.dataset_filters == preset.dataset_filters + assert resolved.include_baseline == preset.include_baseline + assert resolved.scenario_params == preset.scenario_params + + def test_launch_fields_reach_the_run_request(self) -> None: + launch = ResolveScenarioPresetRequest( + target_name="gpt4", + adversarial_target_name="adversary", + initializers=["scorer"], + initializer_args={"scorer": {"threshold": 0.5}}, + labels={"operator": "red"}, + ) + + resolved = ScenarioPresetService.resolve_run_request(preset=_preset(), launch=launch) + + assert resolved.target_name == "gpt4" + assert resolved.adversarial_target_name == "adversary" + assert resolved.initializers == ["scorer"] + assert resolved.initializer_args == {"scorer": {"threshold": 0.5}} + assert resolved.labels == {"operator": "red"} + + +class TestServiceConfiguration: + """Storage source selection.""" + + def test_the_service_is_a_singleton(self) -> None: + assert get_scenario_preset_service() is get_scenario_preset_service() + + def test_configure_source_replaces_the_storage_location(self, tmp_path: Path) -> None: + instance = ScenarioPresetService() + + instance.configure_source(str(tmp_path)) + + assert instance._get_storage().display_source == str(tmp_path) + + def test_an_unconfigured_service_falls_back_to_the_default_directory(self) -> None: + instance = ScenarioPresetService() + + assert instance._get_storage().display_source.endswith("scenario_presets") From 78bcb9439f3acba3f99054a66e775bd39c4b2564 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Fri, 2 Oct 2026 15:57:47 -0400 Subject: [PATCH 07/16] Add scenario preset UI: library, editor, and launch dialog Adds the frontend for scenario presets on top of the preset API: - ScenarioPresetLibrary lists stored presets, surfaces validation issues from the backend, disables Launch for presets this deployment cannot run, and supports delete with confirmation. - ScenarioPresetEditor creates and updates presets, pinning the scenario-owned fields only. Launch-time settings (target, concurrency, retries, labels) are deliberately absent. - LaunchPresetDialog collects the launch-time settings, resolves the preset into a run request, and starts the run. Extracts the scenario-owned configuration out of ScenarioDetail into shared modules (scenarioConfigForm, ScenarioTechniqueSelector, ScenarioDatasetFields, scenarioRunLimits) so the launch form and the preset editor build the same config from one implementation. Reached from the scenario catalog; routed under /scanner/presets. The /edit suffix keeps /scanner/presets/new from shadowing a preset literally named "new". Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- frontend/src/App.tsx | 15 + .../LaunchPresetDialog.styles.ts | 15 + .../LaunchPresetDialog.test.tsx | 252 +++++++++++ .../ScenarioPresets/LaunchPresetDialog.tsx | 186 ++++++++ .../ScenarioPresetEditor.styles.ts | 70 +++ .../ScenarioPresetEditor.test.tsx | 412 +++++++++++++++++ .../ScenarioPresets/ScenarioPresetEditor.tsx | 421 ++++++++++++++++++ .../ScenarioPresetLibrary.styles.ts | 113 +++++ .../ScenarioPresetLibrary.test.tsx | 211 +++++++++ .../ScenarioPresets/ScenarioPresetLibrary.tsx | 326 ++++++++++++++ .../ScenarioPresets/presetRoutes.test.ts | 21 + .../ScenarioPresets/presetRoutes.ts | 16 + .../scenarioPresetForm.test.ts | 213 +++++++++ .../ScenarioPresets/scenarioPresetForm.ts | 103 +++++ .../components/Scenarios/ScenarioCatalog.tsx | 14 +- .../Scenarios/ScenarioDatasetFields.styles.ts | 24 + .../Scenarios/ScenarioDatasetFields.tsx | 97 ++++ .../components/Scenarios/ScenarioDetail.tsx | 390 +++------------- .../ScenarioTechniqueSelector.styles.ts | 60 +++ .../Scenarios/ScenarioTechniqueSelector.tsx | 149 +++++++ .../Scenarios/scenarioConfigForm.ts | 213 +++++++++ .../components/Scenarios/scenarioRunLimits.ts | 24 + frontend/src/services/api.ts | 53 +++ frontend/src/test-utils/scenarioFixtures.ts | 48 ++ frontend/src/types/index.ts | 55 +++ 25 files changed, 3164 insertions(+), 337 deletions(-) create mode 100644 frontend/src/components/ScenarioPresets/LaunchPresetDialog.styles.ts create mode 100644 frontend/src/components/ScenarioPresets/LaunchPresetDialog.test.tsx create mode 100644 frontend/src/components/ScenarioPresets/LaunchPresetDialog.tsx create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetEditor.styles.ts create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetEditor.test.tsx create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetEditor.tsx create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetLibrary.styles.ts create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetLibrary.test.tsx create mode 100644 frontend/src/components/ScenarioPresets/ScenarioPresetLibrary.tsx create mode 100644 frontend/src/components/ScenarioPresets/presetRoutes.test.ts create mode 100644 frontend/src/components/ScenarioPresets/presetRoutes.ts create mode 100644 frontend/src/components/ScenarioPresets/scenarioPresetForm.test.ts create mode 100644 frontend/src/components/ScenarioPresets/scenarioPresetForm.ts create mode 100644 frontend/src/components/Scenarios/ScenarioDatasetFields.styles.ts create mode 100644 frontend/src/components/Scenarios/ScenarioDatasetFields.tsx create mode 100644 frontend/src/components/Scenarios/ScenarioTechniqueSelector.styles.ts create mode 100644 frontend/src/components/Scenarios/ScenarioTechniqueSelector.tsx create mode 100644 frontend/src/components/Scenarios/scenarioConfigForm.ts create mode 100644 frontend/src/components/Scenarios/scenarioRunLimits.ts create mode 100644 frontend/src/test-utils/scenarioFixtures.ts diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index c0eb76ceda..e79629373c 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -21,6 +21,8 @@ import ScenarioHistory from './components/History/ScenarioHistory' import ScenarioCatalog from './components/Scenarios/ScenarioCatalog' import ScenarioDetail from './components/Scenarios/ScenarioDetail' import ScenarioRunPage from './components/Scenarios/ScenarioRunPage' +import ScenarioPresetEditor from './components/ScenarioPresets/ScenarioPresetEditor' +import ScenarioPresetLibrary from './components/ScenarioPresets/ScenarioPresetLibrary' import FeedbackDialog from './components/Feedback/FeedbackDialog' import type { HistoryFilters } from './components/History/historyFilters' import { ConnectionBanner } from './components/ConnectionBanner' @@ -674,6 +676,19 @@ function AppContent({ operatorAlias }: { operatorAlias: string | null }) { } /> } /> + + } + /> + } /> + } /> ({ + scenarioPresetsApi: { + resolve: jest.fn(), + }, + scenariosApi: { + startRun: jest.fn(), + }, +})) + +const mockResolve = scenarioPresetsApi.resolve as jest.Mock +const mockStartRun = scenariosApi.startRun as jest.Mock + +const mockNavigate = jest.fn() + +jest.mock('react-router', () => ({ + ...jest.requireActual('react-router'), + useNavigate: () => mockNavigate, +})) + +const PRESET: ScenarioPreset = { + name: 'nightly_probe', + scenario_name: 'foundry.red_team_agent', + techniques: ['crescendo'], +} + +const OBJECTIVE_TARGET = makeTarget({ target_registry_name: 'gpt4o' }) +const ADVERSARIAL_TARGET: TargetInstance = makeTarget({ + target_registry_name: 'adversary', + capabilities: { + supports_multi_turn: true, + supports_json_schema: false, + supports_json_output: false, + supports_system_prompt: true, + supported_input_modalities: ['text'], + supported_output_modalities: ['text'], + }, +}) + +interface RenderOptions { + defaultObjectiveTarget?: TargetInstance | null + defaultAdversarialTarget?: TargetInstance | null + labels?: Record +} + +const onDismiss = jest.fn() + +function renderDialog({ + defaultObjectiveTarget = OBJECTIVE_TARGET, + defaultAdversarialTarget = null, + labels = {}, +}: RenderOptions = {}) { + return render( + + + + + , + ) +} + +beforeEach(() => { + jest.clearAllMocks() + mockResolve.mockResolvedValue({ scenario_name: 'foundry.red_team_agent', target_name: 'gpt4o' }) + mockStartRun.mockResolvedValue({ scenario_result_id: 'run-1' }) +}) + +describe('LaunchPresetDialog', () => { + it('resolves the preset server-side and navigates to the started run', async () => { + const user = userEvent.setup() + renderDialog({ labels: { op: 'nightly' } }) + + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith('nightly_probe', { + target_name: 'gpt4o', + max_concurrency: 10, + max_retries: 0, + labels: { op: 'nightly' }, + })) + expect(mockStartRun).toHaveBeenCalledWith({ + scenario_name: 'foundry.red_team_agent', + target_name: 'gpt4o', + }) + expect(mockNavigate).toHaveBeenCalledWith('/scanner-history/run-1', { + state: { scenarioName: 'foundry.red_team_agent' }, + }) + }) + + it('omits the adversarial target and labels when neither is set', async () => { + const user = userEvent.setup() + renderDialog() + + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith('nightly_probe', { + target_name: 'gpt4o', + max_concurrency: 10, + max_retries: 0, + })) + }) + + it('sends the default adversarial target when one is configured', async () => { + const user = userEvent.setup() + renderDialog({ defaultAdversarialTarget: ADVERSARIAL_TARGET }) + + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ adversarial_target_name: 'adversary' }), + )) + }) + + it('blocks launching until a target is chosen', () => { + renderDialog({ defaultObjectiveTarget: null }) + + expect(screen.getByTestId('confirm-launch-preset')).toBeDisabled() + }) + + it('ignores a default target that this deployment does not have', () => { + renderDialog({ defaultObjectiveTarget: makeTarget({ target_registry_name: 'retired' }) }) + + expect(screen.getByTestId('confirm-launch-preset')).toBeDisabled() + }) + + it('reports a resolve failure and does not start a run', async () => { + const user = userEvent.setup() + mockResolve.mockRejectedValue(new Error('Unknown scenario "foundry.retired".')) + + renderDialog() + + await user.click(screen.getByTestId('confirm-launch-preset')) + + expect(await screen.findByText('Unknown scenario "foundry.retired".')).toBeInTheDocument() + expect(mockStartRun).not.toHaveBeenCalled() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it('re-enables launch after a failure so the operator can retry', async () => { + const user = userEvent.setup() + mockStartRun.mockRejectedValueOnce(new Error('Target is unavailable.')) + + renderDialog() + + await user.click(screen.getByTestId('confirm-launch-preset')) + + expect(await screen.findByText('Target is unavailable.')).toBeInTheDocument() + expect(screen.getByTestId('confirm-launch-preset')).toBeEnabled() + + await user.click(screen.getByTestId('confirm-launch-preset')) + await waitFor(() => expect(mockNavigate).toHaveBeenCalled()) + }) + + it('dismisses without launching', async () => { + const user = userEvent.setup() + renderDialog() + + await user.click(screen.getByRole('button', { name: 'Cancel' })) + + expect(onDismiss).toHaveBeenCalled() + expect(mockResolve).not.toHaveBeenCalled() + }) + + it('dismisses when the dialog itself is closed', async () => { + const user = userEvent.setup() + renderDialog() + + await user.keyboard('{Escape}') + + expect(onDismiss).toHaveBeenCalled() + }) + + it('sends the target the operator picked instead of the default', async () => { + const user = userEvent.setup() + renderDialog() + + await user.selectOptions(screen.getByRole('combobox', { name: 'Target' }), 'adversary') + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ target_name: 'adversary' }), + )) + }) + + it('sends the adversarial target the operator picked', async () => { + const user = userEvent.setup() + renderDialog() + + await user.selectOptions(screen.getByRole('combobox', { name: 'Adversarial Target' }), 'adversary') + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ adversarial_target_name: 'adversary' }), + )) + }) + + it('sends the concurrency and retry limits the operator set', async () => { + const user = userEvent.setup() + renderDialog() + + const concurrency = screen.getByTestId('preset-max-concurrency-input') + await user.clear(concurrency) + await user.type(concurrency, '4') + const retries = screen.getByTestId('preset-max-retries-input') + await user.clear(retries) + await user.type(retries, '2') + + await user.click(screen.getByTestId('confirm-launch-preset')) + + await waitFor(() => expect(mockResolve).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ max_concurrency: 4, max_retries: 2 }), + )) + }) + + it('does not resolve twice while a launch is in flight', async () => { + const user = userEvent.setup() + let release: () => void = () => {} + mockResolve.mockReturnValue(new Promise((resolve) => { + release = () => resolve({ scenario_name: 'foundry.red_team_agent', target_name: 'gpt4o' }) + })) + + renderDialog() + + await user.click(screen.getByTestId('confirm-launch-preset')) + expect(screen.getByTestId('confirm-launch-preset')).toBeDisabled() + + release() + await waitFor(() => expect(mockNavigate).toHaveBeenCalled()) + expect(mockResolve).toHaveBeenCalledTimes(1) + }) +}) diff --git a/frontend/src/components/ScenarioPresets/LaunchPresetDialog.tsx b/frontend/src/components/ScenarioPresets/LaunchPresetDialog.tsx new file mode 100644 index 0000000000..af5f7bce10 --- /dev/null +++ b/frontend/src/components/ScenarioPresets/LaunchPresetDialog.tsx @@ -0,0 +1,186 @@ +import { type FormEvent, useState } from 'react' + +import { + Button, + Dialog, + DialogActions, + DialogBody, + DialogContent, + DialogSurface, + DialogTitle, + Field, + MessageBar, + MessageBarBody, + Text, +} from '@fluentui/react-components' +import { useNavigate } from 'react-router' + +import TargetSelect from '@/components/Config/TargetSelect' +import SingleStepSpinButton from '@/components/Parameters/SingleStepSpinButton' +import { + DEFAULT_MAX_CONCURRENCY, + DEFAULT_MAX_RETRIES, + MAX_MAX_CONCURRENCY, + MAX_MAX_RETRIES, + MIN_MAX_CONCURRENCY, + MIN_MAX_RETRIES, + resolveSpinButtonValue, +} from '@/components/Scenarios/scenarioRunLimits' +import { scenarioPresetsApi, scenariosApi } from '@/services/api' +import { toApiError } from '@/services/errors' +import type { ScenarioPreset, TargetInstance } from '@/types' +import { scenarioRunRoutePath } from '@/utils/routeParams' + +import { useLaunchPresetDialogStyles } from './LaunchPresetDialog.styles' + +interface LaunchPresetDialogProps { + preset: ScenarioPreset + targets: TargetInstance[] + defaultObjectiveTarget: TargetInstance | null + defaultAdversarialTarget: TargetInstance | null + labels: Record + onDismiss: () => void +} + +function initialTargetName( + candidate: TargetInstance | null, + available: TargetInstance[], +): string { + if (!candidate) { + return '' + } + const match = available.some( + (target) => target.target_registry_name === candidate.target_registry_name, + ) + return match ? candidate.target_registry_name : '' +} + +/** + * Collects the launch-owned fields a preset deliberately omits, then asks the + * server to merge them with the stored preset. Resolution stays server-side so + * the browser never reimplements the preset-to-run mapping. + */ +export default function LaunchPresetDialog({ + preset, + targets, + defaultObjectiveTarget, + defaultAdversarialTarget, + labels, + onDismiss, +}: LaunchPresetDialogProps) { + const styles = useLaunchPresetDialogStyles() + const navigate = useNavigate() + const adversarialTargets = targets.filter( + (target) => target.capabilities?.supports_multi_turn === true, + ) + const [targetName, setTargetName] = useState( + () => initialTargetName(defaultObjectiveTarget, targets), + ) + const [adversarialTargetName, setAdversarialTargetName] = useState( + () => initialTargetName(defaultAdversarialTarget, adversarialTargets), + ) + const [maxConcurrency, setMaxConcurrency] = useState(DEFAULT_MAX_CONCURRENCY) + const [maxRetries, setMaxRetries] = useState(DEFAULT_MAX_RETRIES) + const [submitting, setSubmitting] = useState(false) + const [error, setError] = useState(null) + + const handleSubmit = async (event: FormEvent): Promise => { + event.preventDefault() + if (targetName === '' || submitting) { + return + } + setSubmitting(true) + setError(null) + try { + const request = await scenarioPresetsApi.resolve(preset.name, { + target_name: targetName, + ...(adversarialTargetName === '' ? {} : { adversarial_target_name: adversarialTargetName }), + max_concurrency: maxConcurrency, + max_retries: maxRetries, + ...(Object.keys(labels).length > 0 ? { labels } : {}), + }) + const summary = await scenariosApi.startRun(request) + navigate(scenarioRunRoutePath(summary.scenario_result_id), { + state: { scenarioName: preset.scenario_name }, + }) + } catch (err) { + setError(toApiError(err).detail) + setSubmitting(false) + } + } + + return ( + { if (!data.open) onDismiss() }}> + +
+ + Launch {preset.name} + + + Runs {preset.scenario_name} with this preset's saved configuration. + + {error && ( + + {error} + + )} + setTargetName(target?.target_registry_name ?? '')} + label="Target" + hint="The registered target this run attacks." + placeholder="Select a target" + disabled={submitting} + /> + setAdversarialTargetName(target?.target_registry_name ?? '')} + label="Adversarial Target" + hint="Only used by scenarios that generate attacks with a second target." + placeholder="Use server default" + disabled={submitting} + /> + + setMaxConcurrency(resolveSpinButtonValue(data, maxConcurrency))} + data-testid="preset-max-concurrency-input" + /> + + + setMaxRetries(resolveSpinButtonValue(data, maxRetries))} + data-testid="preset-max-retries-input" + /> + + + + + + + +
+
+
+ ) +} diff --git a/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.styles.ts b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.styles.ts new file mode 100644 index 0000000000..d48adda1f6 --- /dev/null +++ b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.styles.ts @@ -0,0 +1,70 @@ +import { makeStyles, tokens } from '@fluentui/react-components' + +import { mobileTouchTarget, NARROW_VIEWPORT_QUERY } from '@/styles/touchTargets' +import { WORKSPACE_CANVAS_BACKGROUND } from '@/styles/workspaceBackground' + +export const useScenarioPresetEditorStyles = makeStyles({ + root: { + display: 'flex', + flexDirection: 'column', + height: '100%', + width: '100%', + minWidth: 0, + padding: tokens.spacingVerticalXXL, + overflowX: 'hidden', + overflowY: 'auto', + backgroundColor: WORKSPACE_CANVAS_BACKGROUND, + [NARROW_VIEWPORT_QUERY]: { + padding: `${tokens.spacingVerticalL} ${tokens.spacingHorizontalM}`, + }, + }, + headerText: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalXS, + marginBottom: tokens.spacingVerticalL, + }, + subtitle: { + color: tokens.colorNeutralForeground3, + }, + form: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalL, + maxWidth: '60rem', + }, + section: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalM, + padding: tokens.spacingVerticalL, + borderRadius: tokens.borderRadiusMedium, + border: `${tokens.strokeWidthThin} solid ${tokens.colorNeutralStroke2}`, + backgroundColor: tokens.colorNeutralBackground1, + }, + control: { + maxWidth: '32rem', + }, + dynamicParameters: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalM, + }, + actions: { + display: 'flex', + flexWrap: 'wrap', + gap: tokens.spacingHorizontalS, + alignItems: 'center', + }, + touchTarget: { + ...mobileTouchTarget, + }, + centeredState: { + display: 'flex', + flexDirection: 'column', + alignItems: 'center', + gap: tokens.spacingVerticalM, + padding: tokens.spacingVerticalXXL, + textAlign: 'center', + }, +}) diff --git a/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.test.tsx b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.test.tsx new file mode 100644 index 0000000000..930569b8a7 --- /dev/null +++ b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.test.tsx @@ -0,0 +1,412 @@ +import { render, screen, waitFor } from '@testing-library/react' +import userEvent from '@testing-library/user-event' +import { FluentProvider, webLightTheme } from '@fluentui/react-components' +import { MemoryRouter, Route, Routes } from 'react-router' + +import { scenarioPresetsApi, scenariosApi } from '@/services/api' +import { makeScenario } from '@/test-utils/scenarioFixtures' +import type { ScenarioPreset } from '@/types' + +import ScenarioPresetEditor from './ScenarioPresetEditor' + +jest.mock('@/services/api', () => ({ + scenarioPresetsApi: { + get: jest.fn(), + create: jest.fn(), + update: jest.fn(), + }, + scenariosApi: { + listCatalog: jest.fn(), + getScenario: jest.fn(), + }, +})) + +const mockGet = scenarioPresetsApi.get as jest.Mock +const mockCreate = scenarioPresetsApi.create as jest.Mock +const mockUpdate = scenarioPresetsApi.update as jest.Mock +const mockListCatalog = scenariosApi.listCatalog as jest.Mock +const mockGetScenario = scenariosApi.getScenario as jest.Mock + +const mockNavigate = jest.fn() + +jest.mock('react-router', () => ({ + ...jest.requireActual('react-router'), + useNavigate: () => mockNavigate, +})) + +const SCENARIO = makeScenario() + +const STORED_PRESET: ScenarioPreset = { + name: 'nightly_probe', + scenario_name: 'foundry.red_team_agent', + description: 'Nightly smoke test.', + techniques: ['crescendo'], + include_baseline: false, +} + +function apiError(status: number, detail: string): unknown { + return { + isAxiosError: true, + response: { status, data: { detail } }, + } +} + +function renderCreate() { + return render( + + + + } /> + + + , + ) +} + +function renderEdit(name = 'nightly_probe') { + return render( + + + + } + /> + + + , + ) +} + +beforeEach(() => { + jest.clearAllMocks() + mockListCatalog.mockResolvedValue({ + items: [SCENARIO], + pagination: { has_more: false, next_cursor: null }, + }) + mockGetScenario.mockResolvedValue(SCENARIO) + mockGet.mockResolvedValue({ preset: STORED_PRESET, version: 'v1', issues: [] }) + mockCreate.mockResolvedValue({ preset: STORED_PRESET, version: 'v1', issues: [] }) + mockUpdate.mockResolvedValue({ preset: STORED_PRESET, version: 'v2', issues: [] }) +}) + +describe('ScenarioPresetEditor create mode', () => { + it('cannot save before a scenario is chosen', async () => { + renderCreate() + + expect(await screen.findByTestId('scenario-preset-editor')).toBeInTheDocument() + expect(screen.getByTestId('save-preset-btn')).toBeDisabled() + }) + + it('creates a preset from the chosen scenario and name', async () => { + const user = userEvent.setup() + renderCreate() + + await user.type(await screen.findByTestId('preset-name-input'), 'nightly_probe') + await user.click(screen.getByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + await waitFor(() => expect(screen.getByTestId('save-preset-btn')).toBeEnabled()) + await user.click(screen.getByTestId('save-preset-btn')) + + await waitFor(() => expect(mockCreate).toHaveBeenCalledWith({ + name: 'nightly_probe', + scenario_name: 'foundry.red_team_agent', + techniques: ['default_technique'], + include_baseline: true, + })) + expect(mockUpdate).not.toHaveBeenCalled() + expect(mockNavigate).toHaveBeenCalledWith('/scanner/presets') + }) + + it('rejects a name the server pattern would reject, without calling the API', async () => { + const user = userEvent.setup() + renderCreate() + + await user.type(await screen.findByTestId('preset-name-input'), 'Nightly-Probe') + await user.click(screen.getByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + await waitFor(() => expect(screen.getByTestId('save-preset-btn')).toBeEnabled()) + await user.click(screen.getByTestId('save-preset-btn')) + + expect(await screen.findByText(/Use lowercase letters/)).toBeInTheDocument() + expect(mockCreate).not.toHaveBeenCalled() + }) + + it('reports a duplicate name returned by the server', async () => { + const user = userEvent.setup() + mockCreate.mockRejectedValue(apiError(409, 'A preset named "nightly_probe" already exists.')) + renderCreate() + + await user.type(await screen.findByTestId('preset-name-input'), 'nightly_probe') + await user.click(screen.getByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + await waitFor(() => expect(screen.getByTestId('save-preset-btn')).toBeEnabled()) + await user.click(screen.getByTestId('save-preset-btn')) + + expect( + await screen.findByText('A preset named "nightly_probe" already exists.'), + ).toBeInTheDocument() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it('reports a missing admin permission', async () => { + const user = userEvent.setup() + mockCreate.mockRejectedValue(apiError(403, 'Admin access required.')) + renderCreate() + + await user.type(await screen.findByTestId('preset-name-input'), 'nightly_probe') + await user.click(screen.getByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + await waitFor(() => expect(screen.getByTestId('save-preset-btn')).toBeEnabled()) + await user.click(screen.getByTestId('save-preset-btn')) + + expect(await screen.findByText('Admin access required.')).toBeInTheDocument() + }) + + it('reports an unselectable technique set rather than saving an empty one', async () => { + const user = userEvent.setup() + mockGetScenario.mockResolvedValue(makeScenario({ + default_techniques: [], + default_technique: 'all', + aggregate_techniques: ['all'], + })) + renderCreate() + + await user.type(await screen.findByTestId('preset-name-input'), 'nightly_probe') + await user.click(screen.getByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + await waitFor(() => expect(screen.getByTestId('save-preset-btn')).toBeEnabled()) + await user.click(screen.getByTestId('save-preset-btn')) + + expect(await screen.findByText('Select at least one technique.')).toBeInTheDocument() + expect(mockCreate).not.toHaveBeenCalled() + }) + + it('reports a failure to load the scenario the operator picked', async () => { + const user = userEvent.setup() + mockGetScenario.mockRejectedValue(apiError(503, 'Scenario registry is unavailable.')) + renderCreate() + + await user.click(await screen.findByTestId('preset-scenario-select')) + await user.click(await screen.findByRole('option', { name: 'foundry.red_team_agent' })) + + expect(await screen.findByText('Scenario registry is unavailable.')).toBeInTheDocument() + expect(screen.getByTestId('save-preset-btn')).toBeDisabled() + }) +}) + +describe('ScenarioPresetEditor edit mode', () => { + it('loads the stored preset and pins its name and scenario', async () => { + renderEdit() + + expect(await screen.findByTestId('scenario-preset-editor')).toBeInTheDocument() + expect(screen.getByTestId('preset-name-input')).toBeDisabled() + expect(screen.getByTestId('preset-name-input')).toHaveValue('nightly_probe') + expect(screen.getByTestId('preset-description-input')).toHaveValue('Nightly smoke test.') + }) + + it('updates with the version it read so a concurrent edit is not overwritten', async () => { + const user = userEvent.setup() + renderEdit() + + await user.click(await screen.findByTestId('save-preset-btn')) + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ name: 'nightly_probe', techniques: ['crescendo'] }), + 'v1', + )) + expect(mockCreate).not.toHaveBeenCalled() + expect(mockNavigate).toHaveBeenCalledWith('/scanner/presets') + }) + + it('reports a version conflict instead of navigating away', async () => { + const user = userEvent.setup() + mockUpdate.mockRejectedValue(apiError(409, 'The preset changed since it was loaded.')) + renderEdit() + + await user.click(await screen.findByTestId('save-preset-btn')) + + expect(await screen.findByText('The preset changed since it was loaded.')).toBeInTheDocument() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it('warns that pinned techniques this deployment lacks will be dropped', async () => { + mockGet.mockResolvedValue({ + preset: { ...STORED_PRESET, techniques: ['crescendo', 'retired_attack'] }, + version: 'v1', + issues: [], + }) + + renderEdit() + + expect(await screen.findByTestId('dropped-techniques-warning')).toHaveTextContent( + 'retired_attack', + ) + }) + + it('reports a missing preset as not found', async () => { + mockGet.mockRejectedValue(apiError(404, 'No such preset.')) + + renderEdit() + + expect(await screen.findByTestId('editor-error-state')).toHaveTextContent( + 'No preset named "nightly_probe"', + ) + }) + + it('reports an unavailable scenario as unavailable, not as a missing preset', async () => { + mockGetScenario.mockRejectedValue(apiError(404, 'Unknown scenario "foundry.red_team_agent".')) + + renderEdit() + + expect(await screen.findByTestId('scenario-unavailable-warning')).toBeInTheDocument() + expect(screen.queryByTestId('editor-error-state')).not.toBeInTheDocument() + expect(screen.getByTestId('save-preset-btn')).toBeDisabled() + }) + + it('decodes a preset name that was escaped into the route', async () => { + mockGet.mockResolvedValue({ + preset: { ...STORED_PRESET, name: 'a_b' }, + version: 'v1', + issues: [], + }) + + renderEdit('a_b') + + await waitFor(() => expect(mockGet).toHaveBeenCalledWith('a_b')) + }) + + it('saves the edited description', async () => { + const user = userEvent.setup() + renderEdit() + + const description = await screen.findByTestId('preset-description-input') + await user.clear(description) + await user.type(description, 'Weekly smoke test.') + await user.click(screen.getByTestId('save-preset-btn')) + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ description: 'Weekly smoke test.' }), + 'v1', + )) + }) + + it('saves the techniques and baseline choice the operator changed', async () => { + const user = userEvent.setup() + renderEdit() + + await user.click(await screen.findByTestId('technique-default_technique')) + await user.click(screen.getByTestId('baseline-checkbox')) + await user.click(screen.getByTestId('save-preset-btn')) + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ + techniques: ['crescendo', 'default_technique'], + include_baseline: true, + }), + 'v1', + )) + }) + + it('saves the dataset overrides the operator entered', async () => { + const user = userEvent.setup() + renderEdit() + + await user.type(await screen.findByTestId('dataset-override-input'), 'harmbench, xstest') + await user.type(screen.getByTestId('max-dataset-size-input'), '25') + await user.type(screen.getByTestId('harm-categories-filter-input'), 'violence') + await user.type(screen.getByTestId('data-types-filter-input'), 'text') + await user.click(screen.getByTestId('save-preset-btn')) + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ + dataset_names: ['harmbench', 'xstest'], + max_dataset_size: 25, + dataset_filters: { harm_categories: ['violence'], data_types: ['text'] }, + }), + 'v1', + )) + }) + + it('rejects a non-positive dataset size before calling the server', async () => { + const user = userEvent.setup() + renderEdit() + + await user.type(await screen.findByTestId('max-dataset-size-input'), '0') + await user.click(screen.getByTestId('save-preset-btn')) + + expect(await screen.findByText('Max dataset size must be a positive integer.')).toBeInTheDocument() + expect(mockUpdate).not.toHaveBeenCalled() + }) + + it('leaves the editor without saving when cancelled', async () => { + const user = userEvent.setup() + renderEdit() + + await user.click(await screen.findByRole('button', { name: 'Cancel' })) + + expect(mockUpdate).not.toHaveBeenCalled() + expect(mockNavigate).toHaveBeenCalledWith('/scanner/presets') + }) + + it('returns to the library from the not-found state', async () => { + const user = userEvent.setup() + mockGet.mockRejectedValue(apiError(404, 'No such preset.')) + + renderEdit() + + await user.click(await screen.findByRole('button', { name: 'Back to presets' })) + + expect(mockNavigate).toHaveBeenCalledWith('/scanner/presets') + }) + + it('reports a catalog failure that is not a missing preset as an error', async () => { + mockListCatalog.mockRejectedValue(apiError(503, 'Preset storage is not configured.')) + + renderEdit() + + expect(await screen.findByTestId('editor-error-state')).toHaveTextContent( + 'Preset storage is not configured.', + ) + }) +}) + +describe('ScenarioPresetEditor dynamic parameters', () => { + const SCENARIO_WITH_PARAM = makeScenario({ + supported_parameters: [ + { + name: 'max_turns', + type_name: 'int', + required: false, + default: '5', + description: 'Turn budget.', + }, + ], + }) + + it('stores a scenario-specific parameter the operator changed', async () => { + const user = userEvent.setup() + mockGetScenario.mockResolvedValue(SCENARIO_WITH_PARAM) + renderEdit() + + const field = await screen.findByTestId('preset-param-max_turns') + await user.clear(field) + await user.type(field, '9') + await user.click(screen.getByTestId('save-preset-btn')) + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledWith( + 'nightly_probe', + expect.objectContaining({ scenario_params: { max_turns: 9 } }), + 'v1', + )) + }) +}) diff --git a/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.tsx b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.tsx new file mode 100644 index 0000000000..03a5187b18 --- /dev/null +++ b/frontend/src/components/ScenarioPresets/ScenarioPresetEditor.tsx @@ -0,0 +1,421 @@ +import { type FormEvent, useCallback, useEffect, useMemo, useState } from 'react' + +import { + Button, + Field, + Input, + MessageBar, + MessageBarBody, + Option, + Spinner, + Combobox, + Text, + Textarea, +} from '@fluentui/react-components' +import { useNavigate, useParams } from 'react-router' + +import ParameterField from '@/components/Parameters/ParameterField' +import type { ParameterFormValue } from '@/components/Parameters/parameterForm' +import ScenarioDatasetFields from '@/components/Scenarios/ScenarioDatasetFields' +import ScenarioTechniqueSelector from '@/components/Scenarios/ScenarioTechniqueSelector' +import { + buildScenarioConfig, + defaultMaxDatasetSize, + dynamicScenarioParameters, + initialScenarioConfigState, + uniqueTechniqueOptions, + type ScenarioConfigFormState, +} from '@/components/Scenarios/scenarioConfigForm' +import { scenarioPresetsApi, scenariosApi } from '@/services/api' +import { toApiError } from '@/services/errors' +import type { RegisteredScenario, ScenarioPreset } from '@/types' +import { fetchAllPages } from '@/utils/fetchAllPages' +import { routerPathParamValue } from '@/utils/routeParams' + +import { useScenarioPresetEditorStyles } from './ScenarioPresetEditor.styles' +import { PRESETS_ROUTE } from './presetRoutes' +import { + configToPreset, + presetToConfigState, + unknownPresetTechniques, + validatePresetName, +} from './scenarioPresetForm' + +/** Items requested per catalog page while paging the scenario picker's options. */ +const CATALOG_PAGE_SIZE = 200 + +type LoadStatus = 'loading' | 'ready' | 'not-found' | 'error' + +interface LoadedPreset { + preset: ScenarioPreset + version: string +} + +export default function ScenarioPresetEditor({ mode }: { mode: 'create' | 'edit' }) { + const { presetName } = useParams<{ presetName: string }>() + // Remount on navigation between two presets so every field resets to the + // newly loaded document rather than retaining the previous one's edits. + return +} + +interface ScenarioPresetEditorContentProps { + mode: 'create' | 'edit' + presetName: string | undefined +} + +function ScenarioPresetEditorContent({ mode, presetName }: ScenarioPresetEditorContentProps) { + const styles = useScenarioPresetEditorStyles() + const navigate = useNavigate() + const decodedName = presetName === undefined ? '' : routerPathParamValue(presetName) + + const [status, setStatus] = useState('loading') + const [loadError, setLoadError] = useState(null) + const [catalog, setCatalog] = useState([]) + const [loaded, setLoaded] = useState(null) + + const [name, setName] = useState(mode === 'edit' ? decodedName : '') + const [description, setDescription] = useState('') + const [scenario, setScenario] = useState(null) + const [scenarioLoading, setScenarioLoading] = useState(false) + const [config, setConfig] = useState(null) + const [droppedTechniques, setDroppedTechniques] = useState([]) + const [scenarioUnavailable, setScenarioUnavailable] = useState(null) + + const [saving, setSaving] = useState(false) + const [validationError, setValidationError] = useState(null) + const [saveError, setSaveError] = useState(null) + + useEffect(() => { + let cancelled = false + + const load = async (): Promise => { + try { + const scenarios = await fetchAllPages( + (cursor) => scenariosApi.listCatalog(CATALOG_PAGE_SIZE, cursor, false), + undefined, + (entry) => entry.scenario_name, + ) + if (cancelled) return + setCatalog(scenarios) + + if (mode === 'create') { + setStatus('ready') + return + } + + const response = await scenarioPresetsApi.get(decodedName) + if (cancelled) return + setLoaded({ preset: response.preset, version: response.version }) + setName(response.preset.name) + setDescription(response.preset.description ?? '') + setStatus('ready') + + // The list endpoint omits run-size estimates, so the preset's scenario is + // re-fetched in full to get the configured dataset caps the form shows. + // A preset may legitimately name a scenario this deployment lacks; that + // is an unavailable scenario, not a missing preset, so it must not be + // reported as a 404 on the preset itself. + try { + const full = await scenariosApi.getScenario(response.preset.scenario_name) + if (cancelled) return + setScenario(full) + setConfig(presetToConfigState(full, response.preset)) + setDroppedTechniques(unknownPresetTechniques(full, response.preset)) + } catch (scenarioErr) { + if (cancelled) return + setScenarioUnavailable(toApiError(scenarioErr).detail) + } + } catch (err) { + if (cancelled) return + const apiError = toApiError(err) + setLoadError(apiError.detail) + setStatus(apiError.status === 404 ? 'not-found' : 'error') + } + } + + void load() + return () => { + cancelled = true + } + }, [decodedName, mode]) + + const handleScenarioChange = useCallback(async (nextScenarioName: string): Promise => { + setScenarioLoading(true) + setSaveError(null) + setValidationError(null) + setDroppedTechniques([]) + try { + const full = await scenariosApi.getScenario(nextScenarioName) + setScenario(full) + setConfig(initialScenarioConfigState(full)) + } catch (err) { + setSaveError(toApiError(err).detail) + } finally { + setScenarioLoading(false) + } + }, []) + + const dynamicParameters = useMemo( + () => (scenario ? dynamicScenarioParameters(scenario) : []), + [scenario], + ) + const techniqueOptions = useMemo( + () => (scenario ? uniqueTechniqueOptions(scenario).techniques : []), + [scenario], + ) + + const updateConfig = (patch: Partial): void => { + setConfig((current) => (current ? { ...current, ...patch } : current)) + setValidationError(null) + } + + const updateScenarioParam = useCallback((parameterName: string, value: ParameterFormValue) => { + setConfig((current) => (current + ? { ...current, scenarioParamValues: { ...current.scenarioParamValues, [parameterName]: value } } + : current)) + setValidationError(null) + }, []) + + const handleSubmit = async (event: FormEvent): Promise => { + event.preventDefault() + if (saving) { + return + } + setSaveError(null) + + const trimmedName = name.trim() + const nameError = validatePresetName(trimmedName) + if (nameError) { + setValidationError(nameError) + return + } + if (!scenario || !config) { + setValidationError('Select a scenario.') + return + } + const built = buildScenarioConfig({ + techniques: config.techniques, + dynamicParameters, + scenarioParamValues: config.scenarioParamValues, + datasetOverride: config.datasetOverride, + maxDatasetSize: config.maxDatasetSize, + harmCategoriesFilter: config.harmCategoriesFilter, + dataTypesFilter: config.dataTypesFilter, + includeBaseline: config.includeBaseline, + }) + if (!built.ok) { + setValidationError(built.error) + return + } + setValidationError(null) + + const preset = configToPreset( + { name: trimmedName, scenarioName: scenario.scenario_name, description }, + built.config, + ) + + setSaving(true) + try { + if (loaded) { + await scenarioPresetsApi.update(loaded.preset.name, preset, loaded.version) + } else { + await scenarioPresetsApi.create(preset) + } + navigate(PRESETS_ROUTE) + } catch (err) { + setSaveError(toApiError(err).detail) + setSaving(false) + } + } + + if (status === 'loading') { + return ( +
+
+ +
+
+ ) + } + + if (status !== 'ready') { + return ( +
+
+ + + {status === 'not-found' ? `No preset named "${decodedName}".` : loadError} + + + +
+
+ ) + } + + const editing = loaded !== null + + return ( +
+
+ + {editing ? `Edit ${decodedName}` : 'New preset'} + + + A preset stores what to test. The target, concurrency and retries are chosen at launch. + +
+ +
+ {saveError && ( + + {saveError} + + )} + {validationError && ( + + {validationError} + + )} + {droppedTechniques.length > 0 && ( + + + This preset pins techniques this deployment does not offer + ({droppedTechniques.join(', ')}). Saving will drop them. + + + )} + {scenarioUnavailable && ( + + + This preset targets {loaded?.preset.scenario_name}, which is not available here + ({scenarioUnavailable}). Its configuration cannot be edited in this deployment. + + + )} + +
+ + { + setName(data.value) + setValidationError(null) + }} + data-testid="preset-name-input" + /> + + +