From 1d8d9ed1c5dcf74a15d3d8fa1543f9c2a139a2ac Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 1 Oct 2026 23:06:29 -0700 Subject: [PATCH 1/2] Make scenario dataset sources and limits explicit Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../instructions/scenarios.instructions.md | 43 +- doc/code/scenarios/0_scenarios.ipynb | 12 +- doc/code/scenarios/0_scenarios.py | 12 +- doc/scanner/garak.ipynb | 6 +- doc/scanner/garak.py | 6 +- .../Scenarios/ScenarioDetail.test.tsx | 32 + .../components/Scenarios/ScenarioDetail.tsx | 20 +- frontend/src/types/index.ts | 4 +- .../scenario_configuration_resolver.py | 43 +- .../backend/services/scenario_run_service.py | 4 +- pyrit/backend/services/scenario_service.py | 4 +- pyrit/cli/_cli_args.py | 41 +- pyrit/cli/pyrit_scan.py | 3 +- .../seed_datasets/seed_dataset_provider.py | 25 + pyrit/memory/memory_interface.py | 24 +- pyrit/models/catalog/scenario.py | 25 +- pyrit/models/dataset_limit.py | 36 + .../models/scenario_dataset_size_estimate.py | 6 +- pyrit/scenario/__init__.py | 4 + pyrit/scenario/core/__init__.py | 4 + pyrit/scenario/core/dataset_configuration.py | 623 ++++++++++++------ pyrit/scenario/core/scenario.py | 53 +- .../scenarios/adaptive/text_adaptive.py | 6 +- pyrit/scenario/scenarios/airt/cyber.py | 6 +- pyrit/scenario/scenarios/airt/jailbreak.py | 6 +- pyrit/scenario/scenarios/airt/leakage.py | 6 +- pyrit/scenario/scenarios/airt/multilingual.py | 5 +- pyrit/scenario/scenarios/airt/psychosocial.py | 110 +--- .../scenario/scenarios/airt/rapid_response.py | 25 +- pyrit/scenario/scenarios/airt/scam.py | 6 +- .../scenarios/benchmark/adversarial.py | 32 +- .../scenarios/foundry/red_team_agent.py | 6 +- .../scenarios/garak/_prompt_injection.py | 7 +- pyrit/scenario/scenarios/garak/api_key.py | 20 +- .../scenarios/garak/audio_achilles_heel.py | 7 +- pyrit/scenario/scenarios/garak/divergence.py | 7 +- pyrit/scenario/scenarios/garak/doctor.py | 6 +- pyrit/scenario/scenarios/garak/encoding.py | 18 +- .../scenario/scenarios/garak/exploitation.py | 6 +- pyrit/scenario/scenarios/garak/figstep.py | 8 +- .../scenarios/garak/latent_injection.py | 28 +- .../scenarios/garak/package_hallucination.py | 38 +- .../scenario/scenarios/garak/prompt_inject.py | 31 +- .../garak/system_prompt_extraction.py | 33 +- .../scenario/scenarios/garak/web_injection.py | 16 +- .../test_scenario_configuration_resolver.py | 110 +++- tests/unit/backend/test_scenario_resume.py | 31 +- .../unit/backend/test_scenario_run_service.py | 63 +- tests/unit/backend/test_scenario_service.py | 54 +- tests/unit/cli/test_api_client.py | 26 +- tests/unit/cli/test_cli_args.py | 15 +- tests/unit/cli/test_pyrit_scan.py | 14 + tests/unit/cli/test_pyrit_shell.py | 28 + .../datasets/test_seed_dataset_provider.py | 14 + .../test_interface_seed_prompts.py | 21 +- tests/unit/models/test_scenario_request.py | 43 +- tests/unit/scenario/airt/test_cyber.py | 2 +- tests/unit/scenario/airt/test_jailbreak.py | 6 +- tests/unit/scenario/airt/test_leakage.py | 3 + tests/unit/scenario/airt/test_psychosocial.py | 118 ++-- .../unit/scenario/airt/test_rapid_response.py | 19 +- tests/unit/scenario/airt/test_scam.py | 5 +- .../scenario/benchmark/test_adversarial.py | 66 +- .../core/test_dataset_configuration.py | 106 +-- .../scenario/core/test_dataset_sources.py | 430 ++++++++++++ tests/unit/scenario/core/test_scenario.py | 70 ++ .../scenario/foundry/test_red_team_agent.py | 5 +- tests/unit/scenario/garak/test_api_key.py | 33 +- tests/unit/scenario/garak/test_divergence.py | 18 +- tests/unit/scenario/garak/test_encoding.py | 7 +- tests/unit/scenario/garak/test_figstep.py | 1 + .../scenario/garak/test_latent_injection.py | 33 +- .../garak/test_package_hallucination.py | 22 +- .../unit/scenario/garak/test_prompt_inject.py | 27 +- .../garak/test_system_prompt_extraction.py | 10 +- .../scenarios/adaptive/test_text_adaptive.py | 31 +- .../test_default_run_size_estimates.py | 68 +- .../scenario/test_generated_objectives.py | 4 +- 78 files changed, 2241 insertions(+), 725 deletions(-) create mode 100644 pyrit/models/dataset_limit.py create mode 100644 tests/unit/scenario/core/test_dataset_sources.py diff --git a/.github/instructions/scenarios.instructions.md b/.github/instructions/scenarios.instructions.md index e9e18680df..db20a2ef8b 100644 --- a/.github/instructions/scenarios.instructions.md +++ b/.github/instructions/scenarios.instructions.md @@ -28,7 +28,7 @@ class MyScenario(Scenario): super().__init__( version=self.VERSION, technique_class=MyTechnique, - default_dataset_config=DatasetConfiguration(dataset_names=["my_dataset"]), + default_dataset_config=DatasetAttackConfiguration(sources=[DatasetSource(name="my_dataset")]), objective_scorer=objective_scorer or self._get_default_objective_scorer(), scenario_result_id=scenario_result_id, ) @@ -68,7 +68,7 @@ def __init__( super().__init__( version=self.VERSION, technique_class=MyTechnique, - default_dataset_config=DatasetConfiguration(dataset_names=["my_dataset"]), + default_dataset_config=DatasetAttackConfiguration(sources=[DatasetSource(name="my_dataset")]), objective_scorer=objective_scorer, ) ``` @@ -113,30 +113,43 @@ Dropping a common input is not silent: `set_params_from_args` rejects any value ## Dataset Loading -Datasets are read from `CentralMemory`. +Datasets are read from `CentralMemory`. New runs call `prepare_async()` before +reading. Reads, discovery, estimates, and resume must never fetch or store datasets. +Resolve parameter-dependent source names before preparation, not during seed reads. ### Basic — named datasets: ```python -DatasetConfiguration( - dataset_names=["airt_hate", "airt_violence"], - max_dataset_size=10, # optional: sample up to N per dataset +DatasetAttackConfiguration( + sources=[DatasetSource(name="airt_hate"), DatasetSource(name="airt_violence")], + max_per_dataset=5, + max_total=10, ) ``` ### Advanced — custom subclass for filtering: ```python -class MyDatasetConfiguration(DatasetConfiguration): - def get_seed_groups(self) -> dict[str, list[SeedGroup]]: - result = super().get_seed_groups() - # Filter by selected techniques via self._scenario_techniques - return filtered_result +class MyDatasetConfiguration(DatasetAttackConfiguration): + def _build_attack_groups(self, seeds: list[Seed]) -> list[AttackSeedGroup]: + return build_custom_groups(seeds) ``` Options: -- `dataset_names` — load by name from memory -- `seed_groups` — pass explicit groups (mutually exclusive with `dataset_names`) -- `max_dataset_size` — cap per dataset -- Override `_load_seed_groups_for_dataset()` for custom loading +- `sources` selects datasets by name. Source `max_size` overrides `max_per_dataset`; + omitted, `None`, empty, and `"default"` mean inherit; `"all"` means no cap. +- All dataset limits use the same rule: omitted, `None`, empty, and `"default"` use the + default; `"all"` removes that limit; a positive integer sets a cap. +- `max_per_dataset=5` is the named-objective default. `max_total="all"` leaves the + combined selection uncapped. Apply source limits before the total limit. +- `fetch=DatasetFetchPolicy.IF_MISSING` prepares absent registered datasets. + `NEVER` requires stored data. A filter miss must never fetch. +- `seed_groups` and `seeds` are inline alternatives. They never use memory or providers. +- Validators run on full filtered populations before any sampling. +- Ingredients must remain complete. Use `sampling_scope="total_only"` on their + dataset configuration; cap the combined attack groups, not each ingredient dataset. +- Use `with_overrides()` instead of reconstructing a subclass or mutating its defaults. +- Keep custom shaping in `_build_attack_groups()` or `_build_groups_by_dataset_async()`. + Reads use `_collect_seeds_for_dataset_async()` and must not fetch. +- `dataset_names`, `max_dataset_size`, `auto_fetch`, and `per_dataset()` are deprecated. ## Technique Enum diff --git a/doc/code/scenarios/0_scenarios.ipynb b/doc/code/scenarios/0_scenarios.ipynb index 715670b689..3116fa7a63 100644 --- a/doc/code/scenarios/0_scenarios.ipynb +++ b/doc/code/scenarios/0_scenarios.ipynb @@ -65,9 +65,10 @@ " Matrix-shaped scenarios delegate to `build_matrix_atomic_attacks(context=...)` in one line.\n", "\n", "3. **Default Dataset**: Pass `default_dataset_config=` to `super().__init__()` to specify the datasets your scenario uses out of the box.\n", - " - Returns a `DatasetConfiguration` with one or more named datasets (e.g., `DatasetConfiguration(dataset_names=[\"my_dataset\"])`)\n", + " - Returns a `DatasetAttackConfiguration` with named sources (e.g., `sources=[DatasetSource(name=\"my_dataset\")]`)\n", " - Users can override this at runtime via `--dataset-names` in the CLI or by passing a custom `dataset_config` programmatically\n", - " - `DatasetAttackConfiguration` selects at most 5 attack groups unless you set `max_dataset_size`; `max_dataset_size=None` uses all groups\n", + " - Named sources select at most 5 attack groups per dataset by default; `max_total` caps the union. For any limit, omitted/`None`/empty/`\"default\"` uses the default; `\"all\"` removes that limit.\n", + " - New runs prepare missing registered datasets once; reads, estimates, and resume never fetch\n", "\n", "4. **Constructor**: Use `@apply_defaults` decorator and call `super().__init__()` with scenario metadata:\n", " - `name`: Descriptive name for your scenario\n", @@ -120,7 +121,8 @@ "source": [ "from pyrit.common import apply_defaults\n", "from pyrit.scenario import (\n", - " DatasetConfiguration,\n", + " DatasetAttackConfiguration,\n", + " DatasetSource,\n", " Scenario,\n", " ScenarioTechnique,\n", ")\n", @@ -166,7 +168,9 @@ " version=self.VERSION,\n", " objective_scorer=self._objective_scorer,\n", " technique_class=MyTechnique,\n", - " default_dataset_config=DatasetConfiguration(dataset_names=[\"dataset_name\"], max_dataset_size=4),\n", + " default_dataset_config=DatasetAttackConfiguration(\n", + " sources=[DatasetSource(name=\"dataset_name\")], max_total=4\n", + " ),\n", " scenario_result_id=scenario_result_id,\n", " )\n", "\n", diff --git a/doc/code/scenarios/0_scenarios.py b/doc/code/scenarios/0_scenarios.py index 927f089d18..2cf481886b 100644 --- a/doc/code/scenarios/0_scenarios.py +++ b/doc/code/scenarios/0_scenarios.py @@ -67,9 +67,10 @@ # Matrix-shaped scenarios delegate to `build_matrix_atomic_attacks(context=...)` in one line. # # 3. **Default Dataset**: Pass `default_dataset_config=` to `super().__init__()` to specify the datasets your scenario uses out of the box. -# - Returns a `DatasetConfiguration` with one or more named datasets (e.g., `DatasetConfiguration(dataset_names=["my_dataset"])`) +# - Returns a `DatasetAttackConfiguration` with named sources (e.g., `sources=[DatasetSource(name="my_dataset")]`) # - Users can override this at runtime via `--dataset-names` in the CLI or by passing a custom `dataset_config` programmatically -# - `DatasetAttackConfiguration` selects at most 5 attack groups unless you set `max_dataset_size`; `max_dataset_size=None` uses all groups +# - Named sources select at most 5 attack groups per dataset by default; `max_total` caps the union. For any limit, omitted/`None`/empty/`"default"` uses the default; `"all"` removes that limit. +# - New runs prepare missing registered datasets once; reads, estimates, and resume never fetch # # 4. **Constructor**: Use `@apply_defaults` decorator and call `super().__init__()` with scenario metadata: # - `name`: Descriptive name for your scenario @@ -97,7 +98,8 @@ # %% from pyrit.common import apply_defaults from pyrit.scenario import ( - DatasetConfiguration, + DatasetAttackConfiguration, + DatasetSource, Scenario, ScenarioTechnique, ) @@ -143,7 +145,9 @@ def __init__( version=self.VERSION, objective_scorer=self._objective_scorer, technique_class=MyTechnique, - default_dataset_config=DatasetConfiguration(dataset_names=["dataset_name"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name="dataset_name")], max_total=4 + ), scenario_result_id=scenario_result_id, ) diff --git a/doc/scanner/garak.ipynb b/doc/scanner/garak.ipynb index 5cfdceba50..890d0bb36f 100644 --- a/doc/scanner/garak.ipynb +++ b/doc/scanner/garak.ipynb @@ -917,7 +917,7 @@ "**Available techniques:** `GetKey` and `CompleteKey`. `DEFAULT` and `ALL` both select the two\n", "techniques. `max_dataset_size` samples across all selected technique populations, not per service.\n", "The base scenario persists the sample for resume. Use `ApiKeyDatasetConfiguration` with\n", - "`max_dataset_size=None` to run all 348 requests. Standard technique converter stacks are supported." + "`max_total=\"all\"` to run all 348 requests. Standard technique converter stacks are supported." ] }, { @@ -1115,7 +1115,7 @@ "groups, shared by six default techniques (552 execution units). Sampling reserves one group\n", "per selected family/trigger pair, then fills the remaining budget without replacement.\n", "A smaller budget than the number of pairs raises an error. An explicit dataset configuration\n", - "with `max_dataset_size=None` uses the complete assembled population. Saved runs replay the sample.\n", + "with `max_total=\"all\"` uses the complete assembled population. Saved runs replay the sample.\n", "\n", "This is not Garak's exact sampling policy: its lightweight probes cap final prompts at 64\n", "per family without guaranteed coverage. PyRIT also applies all selected separators to all\n", @@ -1935,7 +1935,7 @@ "\n", "**Available techniques:** `Repeat`, `DEFAULT`, and `ALL` all select the same probe.\n", "The default budget is 10 prompts across the entire dataset, not per word. Use\n", - "`DivergenceDatasetConfiguration(max_dataset_size=None, dataset_names=[\"garak_divergence\"])`\n", + "`DivergenceDatasetConfiguration(max_per_dataset=\"all\", max_total=\"all\", dataset_names=[\"garak_divergence\"])`\n", "to run all 36 prompts. The example below samples only two." ] }, diff --git a/doc/scanner/garak.py b/doc/scanner/garak.py index a4ab1423b0..1fdbc0857b 100644 --- a/doc/scanner/garak.py +++ b/doc/scanner/garak.py @@ -317,7 +317,7 @@ # **Available techniques:** `GetKey` and `CompleteKey`. `DEFAULT` and `ALL` both select the two # techniques. `max_dataset_size` samples across all selected technique populations, not per service. # The base scenario persists the sample for resume. Use `ApiKeyDatasetConfiguration` with -# `max_dataset_size=None` to run all 348 requests. Standard technique converter stacks are supported. +# `max_total="all"` to run all 348 requests. Standard technique converter stacks are supported. # %% api_key_scenario = ApiKey() @@ -383,7 +383,7 @@ # groups, shared by six default techniques (552 execution units). Sampling reserves one group # per selected family/trigger pair, then fills the remaining budget without replacement. # A smaller budget than the number of pairs raises an error. An explicit dataset configuration -# with `max_dataset_size=None` uses the complete assembled population. Saved runs replay the sample. +# with `max_total="all"` uses the complete assembled population. Saved runs replay the sample. # # This is not Garak's exact sampling policy: its lightweight probes cap final prompts at 64 # per family without guaranteed coverage. PyRIT also applies all selected separators to all @@ -599,7 +599,7 @@ # # **Available techniques:** `Repeat`, `DEFAULT`, and `ALL` all select the same probe. # The default budget is 10 prompts across the entire dataset, not per word. Use -# `DivergenceDatasetConfiguration(max_dataset_size=None, dataset_names=["garak_divergence"])` +# `DivergenceDatasetConfiguration(max_per_dataset="all", max_total="all", dataset_names=["garak_divergence"])` # to run all 36 prompts. The example below samples only two. # %% diff --git a/frontend/src/components/Scenarios/ScenarioDetail.test.tsx b/frontend/src/components/Scenarios/ScenarioDetail.test.tsx index 72834244bf..64581e443e 100644 --- a/frontend/src/components/Scenarios/ScenarioDetail.test.tsx +++ b/frontend/src/components/Scenarios/ScenarioDetail.test.tsx @@ -1029,6 +1029,38 @@ describe('ScenarioDetail', () => { expect(mockStartRun.mock.calls[0][0]).not.toHaveProperty('max_dataset_size') }) + it('sends an explicit all total for both estimation and launch', async () => { + const user = userEvent.setup() + renderDetail('/scanner/foundry.red_team_agent') + const unlimited = await screen.findByRole('checkbox', { name: 'No total dataset limit' }) + await user.click(unlimited) + expect(screen.getByRole('spinbutton', { name: 'Max dataset size' })).toBeDisabled() + await waitFor(() => expect(mockEstimateRun.mock.calls.at(-1)?.[1]).toHaveProperty('max_dataset_size', 'all')) + await confirmRunPreview(user) + await waitFor(() => expect(mockStartRun).toHaveBeenCalled()) + expect(mockStartRun.mock.calls[0][0]).toHaveProperty('max_dataset_size', 'all') + }) + + it('restores the scenario default when no-total-limit is unchecked', async () => { + const user = userEvent.setup() + mockGetScenario.mockResolvedValueOnce(makeScenario({ + default_run_size: { + ...makeEstimate(10), + dataset_limit: { state: 'value', value: 5 }, + }, + })) + renderDetail('/scanner/foundry.red_team_agent') + const unlimited = await screen.findByRole('checkbox', { name: 'No total dataset limit' }) + await user.click(unlimited) + await waitFor(() => expect(mockEstimateRun.mock.calls.at(-1)?.[1]).toHaveProperty('max_dataset_size', 'all')) + await user.click(unlimited) + expect(screen.getByRole('spinbutton', { name: 'Max dataset size' })).toHaveValue(5) + await waitFor(() => expect(mockEstimateRun.mock.calls.at(-1)?.[1]).not.toHaveProperty('max_dataset_size')) + await confirmRunPreview(user) + await waitFor(() => expect(mockStartRun).toHaveBeenCalled()) + expect(mockStartRun.mock.calls[0][0]).not.toHaveProperty('max_dataset_size') + }) + it('includes dataset overrides and filters when provided', async () => { const user = userEvent.setup() renderDetail('/scanner/foundry.red_team_agent') diff --git a/frontend/src/components/Scenarios/ScenarioDetail.tsx b/frontend/src/components/Scenarios/ScenarioDetail.tsx index ed534d9d32..6eb4f16fc7 100644 --- a/frontend/src/components/Scenarios/ScenarioDetail.tsx +++ b/frontend/src/components/Scenarios/ScenarioDetail.tsx @@ -310,9 +310,11 @@ function buildEstimateRequest({ scenarioParams = result.parameters } - let maxDatasetSizeValue: number | undefined + let maxDatasetSizeValue: number | 'all' | undefined const trimmedMaxDatasetSize = maxDatasetSize.trim() - if (trimmedMaxDatasetSize.length > 0) { + if (trimmedMaxDatasetSize === 'all') { + maxDatasetSizeValue = 'all' + } else if (trimmedMaxDatasetSize.length > 0) { const parsed = Number(trimmedMaxDatasetSize) if (!Number.isInteger(parsed) || parsed < 1) { return { ok: false, error: 'Max dataset size must be a positive integer.' } @@ -667,6 +669,8 @@ function ScenarioLaunchForm({ : '' const datasetSizeLabel = scenario.default_run_size.dataset_limit.state === 'not_applicable' ? 'Not applicable' + : maxDatasetSize === 'all' + ? 'No total limit' : maxDatasetSize.trim() || configuredDefaultMaxDatasetSize || 'Scenario default' const estimateResult = useMemo( () => buildEstimateRequest({ @@ -1144,12 +1148,20 @@ function ScenarioLaunchForm({ className={styles.numberInput} type="number" min={1} - value={maxDatasetSize} - disabled={submitting || scenario.default_run_size.dataset_limit.state === 'not_applicable'} + value={maxDatasetSize === 'all' ? '' : maxDatasetSize} + disabled={submitting || maxDatasetSize === 'all' + || scenario.default_run_size.dataset_limit.state === 'not_applicable'} onChange={(_, data) => setMaxDatasetSize(data.value)} data-testid="max-dataset-size-input" /> + setMaxDatasetSize(data.checked === true ? 'all' : configuredDefaultMaxDatasetSize)} + /> + Per-dataset limits still apply. An empty input uses the scenario default total. | null max_concurrency?: number max_retries?: number @@ -770,7 +770,7 @@ export interface ScenarioRunSizeEstimateRequest { adversarial_target_name?: string | null techniques?: string[] | null dataset_names?: string[] | null - max_dataset_size?: number | null + max_dataset_size?: number | 'all' | 'default' | '' | null dataset_filters?: Record | null include_baseline?: boolean | null scenario_params?: Record | null diff --git a/pyrit/backend/services/scenario_configuration_resolver.py b/pyrit/backend/services/scenario_configuration_resolver.py index 9b679c6fd2..a323fb0bff 100644 --- a/pyrit/backend/services/scenario_configuration_resolver.py +++ b/pyrit/backend/services/scenario_configuration_resolver.py @@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any +from pyrit.models.dataset_limit import DatasetLimit, normalize_dataset_limit from pyrit.registry import ConverterRegistry, ScenarioRegistry, TargetRegistry from pyrit.scenario.core.scenario_target_defaults import validate_default_adversarial_target @@ -90,7 +91,7 @@ def resolve_configuration( objective_target: Any | None = None, techniques: list[str] | None = None, dataset_names: list[str] | None = None, - max_dataset_size: int | None = None, + max_dataset_size: DatasetLimit = "default", dataset_filters: dict[str, list[str]] | None = None, include_baseline: bool | None = None, max_concurrency: int | None = None, @@ -100,6 +101,8 @@ def resolve_configuration( """ Resolve shared launch and estimate fields into scenario parameters. + Omitted, null, empty, and "default" limits retain the default. "all" removes only the total limit. + Returns: dict[str, Any]: Values accepted by ``Scenario.set_params_from_args``. @@ -119,7 +122,9 @@ def resolve_configuration( resolved["memory_labels"] = memory_labels filters = dataset_filters or {} - needs_introspection = bool(techniques) or bool(dataset_names) or max_dataset_size is not None or bool(filters) + total_limit = normalize_dataset_limit(max_dataset_size) + has_total_override = total_limit != "default" + needs_introspection = bool(techniques) or bool(dataset_names) or has_total_override or bool(filters) if not needs_introspection: return resolved @@ -141,29 +146,21 @@ def resolve_configuration( if technique_converters: resolved["technique_converters"] = technique_converters - if dataset_names or max_dataset_size is not None or filters: + if dataset_names or has_total_override or filters: + from pyrit.scenario import DatasetSource + default_config = introspection_instance._default_dataset_config + config = default_config.with_overrides(filters=filters) if dataset_names: - default_config_class = type(default_config) - try: - resolved["dataset_config"] = default_config_class( - dataset_names=dataset_names, - max_dataset_size=( - max_dataset_size if max_dataset_size is not None else default_config.max_dataset_size - ), - filters=filters or None, - ) - except TypeError as exc: - raise ValueError( - f"Scenario '{scenario_name}' does not support overriding dataset names through " - f"its {default_config_class.__name__} configuration: {exc}" - ) from exc - else: - if max_dataset_size is not None: - default_config.max_dataset_size = max_dataset_size - if filters: - default_config.update_filters(filters=filters) - resolved["dataset_config"] = default_config + existing_sources = {source.name: source for source in config.sources} + config = config.with_overrides( + sources=[ + existing_sources[name] if name in existing_sources else DatasetSource(name=name) + for name in dataset_names + ] + ) + config = config.with_overrides(max_total=total_limit) + resolved["dataset_config"] = config return resolved diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index 50808b204c..19849c5dca 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -326,13 +326,13 @@ def _restore_launch_request(self, *, stored: ScenarioResult) -> RunScenarioReque raw_request = {"adversarial_target_name": None, **raw_request} if ( not isinstance(raw_request, dict) - or any(name not in raw_request for name in _LAUNCH_REQUEST_FIELDS) + or any(name not in raw_request for name in _LAUNCH_REQUEST_FIELDS if name != "max_dataset_size") or raw_request["include_baseline"] is None ): raise ScenarioRunConflictError("The saved launch configuration is incomplete; resume was not started.") try: request = RunScenarioRequest.model_validate( - {name: raw_request[name] for name in _LAUNCH_REQUEST_FIELDS}, strict=True + {name: raw_request[name] for name in _LAUNCH_REQUEST_FIELDS if name in raw_request}, strict=True ) except ValidationError as exc: raise ScenarioRunConflictError( diff --git a/pyrit/backend/services/scenario_service.py b/pyrit/backend/services/scenario_service.py index 1596e2c59b..443b02e5a5 100644 --- a/pyrit/backend/services/scenario_service.py +++ b/pyrit/backend/services/scenario_service.py @@ -22,7 +22,6 @@ ) from pyrit.registry import ScenarioMetadata, ScenarioRegistry from pyrit.scenario.core import Scenario, override_default_adversarial_target -from pyrit.scenario.core.dataset_configuration import read_only_dataset_resolution logger = logging.getLogger(__name__) _ESTIMATE_CACHE_SIZE = 128 @@ -408,8 +407,7 @@ async def _run_default_estimate_async( construction_complete.set() if execution_timed_out.is_set(): raise asyncio.CancelledError - with read_only_dataset_resolution(): - return await scenario.get_default_run_size_estimate_async() + return await scenario.get_default_run_size_estimate_async() def _clear_estimate_task(self, *, task: _EstimateTask, cache_key: _EstimateCacheKey) -> None: """Remove a completed single-flight task without disturbing a replacement.""" diff --git a/pyrit/cli/_cli_args.py b/pyrit/cli/_cli_args.py index 1556feda10..3d92084889 100644 --- a/pyrit/cli/_cli_args.py +++ b/pyrit/cli/_cli_args.py @@ -21,7 +21,7 @@ import shlex from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Any, get_args, get_origin +from typing import TYPE_CHECKING, Any, Literal, get_args, get_origin from pyrit.common.cli_helpers import ( CONFIG_FILE_HELP, @@ -158,6 +158,21 @@ def validate_integer(value: str, *, name: str = "value", min_value: int | None = return int_value +def parse_dataset_limit(value: str) -> int | Literal["all", "default"]: + """ + Parse a positive total limit, ``default``, or ``all``. + + Returns: + int | Literal["all", "default"]: The requested total limit. + + Raises: + ValueError: If the value is not a supported dataset limit. + """ + from pyrit.models.dataset_limit import normalize_dataset_limit + + return normalize_dataset_limit(value) + + # --------------------------------------------------------------------------- # Argparse adapter # --------------------------------------------------------------------------- @@ -412,9 +427,9 @@ def _coerce_filter_values(value: str) -> list[str]: "database": "Database type to use for memory storage", "log_level": "Logging level", "dataset_names": "List of dataset names to use instead of scenario defaults (e.g., harmbench advbench). " - "Creates a new dataset config; fetches all items unless --max-dataset-size is also specified", - "max_dataset_size": "Maximum number of items to use from the dataset (must be >= 1). " - "Limits new datasets if --dataset-names provided, otherwise overrides scenario's default limit", + "Retains source settings for names already selected by the scenario", + "max_dataset_size": "Total attack-group limit (positive integer, 'default', or 'all' for no total limit). " + "Omit to retain the scenario default. Per-dataset limits still apply", "dataset_filters": "Dataset seed filters as KEY=VALUE tokens " "(e.g., harm_categories=cyber data_types=text). Keys filter seeds before sizing. " "List values may be comma-separated, but semantics differ per key: " @@ -625,7 +640,7 @@ class _ArgSpec: _MAX_DATASET_SIZE_ARG = _ArgSpec( flags=["--max-dataset-size"], result_key="max_dataset_size", - parser=lambda v: validate_integer(v, name="--max-dataset-size", min_value=1), + parser=parse_dataset_limit, ) _DATASET_FILTERS_ARG = _ArgSpec( flags=["--dataset-filters"], @@ -674,8 +689,7 @@ def _parse_shell_arguments(*, parts: list[str], arg_specs: list[_ArgSpec]) -> di arg_specs: Argument specifications that this command accepts. Returns: - Dictionary mapping each spec's ``result_key`` to its parsed value, - defaulting to ``None`` for arguments not present in *parts*. + Dictionary mapping supplied arguments to their parsed values. Absent arguments are omitted. Raises: ValueError: On unknown flags or missing values. @@ -686,8 +700,7 @@ def _parse_shell_arguments(*, parts: list[str], arg_specs: list[_ArgSpec]) -> di for flag in spec.flags: flag_to_spec[flag] = spec - # Initialise result with None defaults - result: dict[str, Any] = {spec.result_key: None for spec in arg_specs} + result: dict[str, Any] = {} i = 0 while i < len(parts): @@ -731,8 +744,8 @@ def parse_run_arguments(*, args_string: str, declared_params: list[Parameter] | ``scenario__`` in the result dict. Returns: - Dictionary mapping built-in result_keys (and ``scenario__*`` keys for - any declared params) to their parsed values. ``scenario_name`` is + Dictionary mapping supplied built-in result_keys (and ``scenario__*`` keys for + supplied declared params) to their parsed values. ``scenario_name`` is always populated from the first positional token. Raises: @@ -762,9 +775,9 @@ def parse_list_targets_arguments(*, args_string: str) -> dict[str, Any]: args_string: Space-separated argument string (e.g., "--initializers target"). Returns: - Dictionary with parsed arguments: - - initializers: list[str | dict[str, Any]] | None - - initialization_scripts: list[str] | None + Dictionary containing only supplied arguments: + - initializers: list[str | dict[str, Any]] + - initialization_scripts: list[str] Raises: ValueError: If parsing or validation fails. diff --git a/pyrit/cli/pyrit_scan.py b/pyrit/cli/pyrit_scan.py index 1eea719358..f912bb5bb5 100644 --- a/pyrit/cli/pyrit_scan.py +++ b/pyrit/cli/pyrit_scan.py @@ -31,6 +31,7 @@ collapse_dataset_filters, non_negative_int, parse_dataset_filter, + parse_dataset_limit, positive_int, validate_log_level_argparse, ) @@ -311,7 +312,7 @@ def _add_run_arguments(*, parser: ArgumentParser, scenario_params: list[Paramete parser.add_argument("--max-retries", type=non_negative_int, help=ARG_HELP["max_retries"]) parser.add_argument("--memory-labels", type=str, help=ARG_HELP["memory_labels"]) parser.add_argument("--dataset-names", type=str, nargs="+", help=ARG_HELP["dataset_names"]) - parser.add_argument("--max-dataset-size", type=positive_int, help=ARG_HELP["max_dataset_size"]) + parser.add_argument("--max-dataset-size", type=parse_dataset_limit, help=ARG_HELP["max_dataset_size"]) parser.add_argument( "--dataset-filters", type=parse_dataset_filter, diff --git a/pyrit/datasets/seed_datasets/seed_dataset_provider.py b/pyrit/datasets/seed_datasets/seed_dataset_provider.py index c5ff12426a..edfffa942b 100644 --- a/pyrit/datasets/seed_datasets/seed_dataset_provider.py +++ b/pyrit/datasets/seed_datasets/seed_dataset_provider.py @@ -107,6 +107,31 @@ def get_all_providers(cls) -> dict[str, type["SeedDatasetProvider"]]: cls._materialize_builtin_providers() return cls._registry.copy() + @classmethod + async def get_providers_by_name_async(cls, *, dataset_names: list[str]) -> dict[str, "SeedDatasetProvider"]: + """ + Resolve registered providers by name without parsing metadata or fetching seeds. + + Returns: + dict[str, SeedDatasetProvider]: Matching providers; missing names are omitted. + """ + return await asyncio.to_thread(cls._get_providers_by_name, dataset_names=dataset_names) + + @classmethod + def _get_providers_by_name(cls, *, dataset_names: list[str]) -> dict[str, "SeedDatasetProvider"]: + cls._materialize_builtin_providers() + requested = set(dataset_names) + providers: dict[str, SeedDatasetProvider] = {} + for provider_class in cls._registry.values(): + provider = provider_class() + name = provider.dataset_name + if name not in requested: + continue + if name in providers: + raise ValueError(f"Multiple registered providers have dataset name '{name}'.") + providers[name] = provider + return providers + @classmethod async def get_all_dataset_names_async(cls, filters: SeedDatasetFilter | None = None) -> list[str]: """ diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 4138ef40a4..1d6e90589b 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -4213,20 +4213,20 @@ def _execute_get_seed_dataset_names(self) -> Sequence[str]: Returns: Sequence[str]: A list of unique dataset names. + + Raises: + SQLAlchemyError: If the dataset-name query fails. """ try: - entries: Sequence[SeedEntry] = self._query_entries( - SeedEntry, - conditions=and_(SeedEntry.dataset_name.isnot(None), SeedEntry.dataset_name != ""), - distinct=True, - ) - # Extract unique dataset names from the entries - dataset_names: set[str] = set() - for entry in entries: - if entry.dataset_name: - dataset_names.add(entry.dataset_name) - return list(dataset_names) - except Exception as e: + with closing(self._get_session()) as session: + statement = ( + select(SeedEntry.dataset_name) + .where(SeedEntry.dataset_name.isnot(None), SeedEntry.dataset_name != "") + .distinct() + ) + names: Sequence[str | None] = session.scalars(statement).all() + return [name for name in names if name is not None] + except SQLAlchemyError as e: logger.exception(f"Failed to retrieve dataset names with error {e}") raise diff --git a/pyrit/models/catalog/scenario.py b/pyrit/models/catalog/scenario.py index b8c234bd47..13f0f5a0af 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -16,10 +16,11 @@ from datetime import datetime from enum import Enum from math import prod -from typing import Any, Literal +from typing import Annotated, Any, Literal -from pydantic import AliasChoices, BaseModel, Field, computed_field, field_validator, model_validator +from pydantic import AliasChoices, BaseModel, BeforeValidator, Field, computed_field, field_validator, model_validator +from pyrit.models.dataset_limit import normalize_dataset_limit from pyrit.models.parameter import Parameter from pyrit.models.results.scenario_result import ScenarioRunState from pyrit.models.retry_event import RetryEvent @@ -46,6 +47,12 @@ DATASET_FILTERS: frozenset[str] = frozenset({"harm_categories", "data_types"}) +_RequestDatasetLimit = Annotated[ + Annotated[int, Field(ge=1, strict=True)] | Literal["all", "default"] | None, + BeforeValidator(normalize_dataset_limit), +] + + def _validate_dataset_filter_mapping( value: dict[str, list[str]] | None, ) -> dict[str, list[str]] | None: @@ -374,7 +381,12 @@ class ScenarioRunSizeEstimateRequest(BaseModel): dataset_names: list[str] | None = Field( None, description="Dataset names to estimate (uses scenario default if omitted)" ) - max_dataset_size: int | None = Field(None, ge=1, description="Maximum selected logical seed groups") + max_dataset_size: _RequestDatasetLimit = Field( + "default", + description="Total selected logical seed-group limit. Omitted, null, empty, or 'default': scenario default. " + "'all': no total limit. " + "Per-dataset limits remain in effect.", + ) dataset_filters: dict[str, list[str]] | None = Field( None, description="Dataset seed filters keyed by field. Accepted keys: harm_categories, data_types.", @@ -415,7 +427,12 @@ 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: _RequestDatasetLimit = Field( + "default", + description="Total selected logical seed-group limit. Omitted, null, empty, or 'default': scenario default. " + "'all': no total limit. " + "Per-dataset limits remain in effect.", + ) dataset_filters: dict[str, list[str]] | None = Field( None, description=( diff --git a/pyrit/models/dataset_limit.py b/pyrit/models/dataset_limit.py new file mode 100644 index 0000000000..c9695e3115 --- /dev/null +++ b/pyrit/models/dataset_limit.py @@ -0,0 +1,36 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Shared input values for scenario dataset limits.""" + +from typing import Literal + +ResolvedDatasetLimit = int | Literal["all"] +DatasetLimit = ResolvedDatasetLimit | Literal["default", ""] | None + + +def normalize_dataset_limit(value: object) -> ResolvedDatasetLimit | Literal["default"]: + """ + Normalize a dataset limit without choosing a scenario-specific default. + + Returns: + ResolvedDatasetLimit | Literal["default"]: A positive cap, all, or default. + + Raises: + ValueError: If the value is not a supported dataset limit. + """ + if value is None: + return "default" + if isinstance(value, str): + value = value.strip().lower() + if value in ("", "default"): + return "default" + if value == "all": + return "all" + try: + value = int(value) + except ValueError: + raise ValueError("Dataset limit must be a positive integer, 'default', or 'all'.") from None + if type(value) is not int or value < 1: + raise ValueError("Dataset limit must be a positive integer, 'default', or 'all'.") + return value diff --git a/pyrit/models/scenario_dataset_size_estimate.py b/pyrit/models/scenario_dataset_size_estimate.py index f6b19d7542..90183990e5 100644 --- a/pyrit/models/scenario_dataset_size_estimate.py +++ b/pyrit/models/scenario_dataset_size_estimate.py @@ -8,6 +8,8 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator +from pyrit.models.dataset_limit import ResolvedDatasetLimit + class ScenarioDatasetSizeEstimateKind(str, Enum): """Meaning of a population budget.""" @@ -47,9 +49,9 @@ class IndeterminateDatasetSize(_BudgetModel): ] -def scenario_dataset_size_from_limit(limit: int | None) -> BoundedDatasetSize | AllAvailableDatasetSize: +def scenario_dataset_size_from_limit(limit: ResolvedDatasetLimit) -> BoundedDatasetSize | AllAvailableDatasetSize: """Return a finite-source budget from an optional selection limit.""" - return AllAvailableDatasetSize() if limit is None else BoundedDatasetSize(value=limit) + return AllAvailableDatasetSize() if limit == "all" else BoundedDatasetSize(value=limit) class DatasetLimitState(str, Enum): diff --git a/pyrit/scenario/__init__.py b/pyrit/scenario/__init__.py index 2f433906bf..16ff045417 100644 --- a/pyrit/scenario/__init__.py +++ b/pyrit/scenario/__init__.py @@ -24,6 +24,8 @@ CompoundDatasetAttackConfiguration, DatasetAttackConfiguration, DatasetConfiguration, + DatasetFetchPolicy, + DatasetSource, DatasetSourceKind, ResolvedDataset, Scenario, @@ -89,6 +91,8 @@ def find_spec( "CompoundDatasetAttackConfiguration": "pyrit.scenario.core.dataset_configuration", "DatasetAttackConfiguration": "pyrit.scenario.core.dataset_configuration", "DatasetConfiguration": "pyrit.scenario.core.dataset_configuration", + "DatasetFetchPolicy": "pyrit.scenario.core.dataset_configuration", + "DatasetSource": "pyrit.scenario.core.dataset_configuration", "DatasetSourceKind": "pyrit.scenario.core.dataset_configuration", "Parameter": "pyrit.models.parameter", "ResolvedDataset": "pyrit.scenario.core.dataset_configuration", diff --git a/pyrit/scenario/core/__init__.py b/pyrit/scenario/core/__init__.py index 992393a906..e17102a84c 100644 --- a/pyrit/scenario/core/__init__.py +++ b/pyrit/scenario/core/__init__.py @@ -24,6 +24,8 @@ DatasetAttackConfiguration, DatasetConfiguration, DatasetConstraintError, + DatasetFetchPolicy, + DatasetSource, DatasetSourceKind, ResolvedDataset, require_nonempty, @@ -45,6 +47,8 @@ "DatasetAttackConfiguration": "pyrit.scenario.core.dataset_configuration", "DatasetConfiguration": "pyrit.scenario.core.dataset_configuration", "DatasetConstraintError": "pyrit.scenario.core.dataset_configuration", + "DatasetFetchPolicy": "pyrit.scenario.core.dataset_configuration", + "DatasetSource": "pyrit.scenario.core.dataset_configuration", "DatasetSourceKind": "pyrit.scenario.core.dataset_configuration", "INLINE_DATASET_NAME": "pyrit.scenario.core.dataset_configuration", "Parameter": "pyrit.models.parameter", diff --git a/pyrit/scenario/core/dataset_configuration.py b/pyrit/scenario/core/dataset_configuration.py index 15fa65c659..13ed15e5ad 100644 --- a/pyrit/scenario/core/dataset_configuration.py +++ b/pyrit/scenario/core/dataset_configuration.py @@ -11,31 +11,30 @@ Constraints are expressed through a single mechanism: ``validators``. Each validator is a ``Callable[[ResolvedDataset], None]`` that raises ``DatasetConstraintError`` on violation. -Validators run against the fully resolved dataset (before ``max_dataset_size`` sampling), +Validators run against the fully resolved dataset (before source or total sampling), so they describe the dataset itself, not the sampled subset. The ``ResolvedDataset`` they receive also carries the ``DatasetSourceKind`` (inline vs from memory) and the contributing ``dataset_names``, which lets a scenario require or forbid inline seeds -- useful for CLI flags such as ``--objectives`` -- restrict which datasets it will resolve from, or require a particular seed type (e.g. ``require_seed_type(SeedObjective)``). -Memory is the source of truth. When a configured dataset name is not yet in memory and -``auto_fetch`` is enabled (the default), the resolver transparently fetches the dataset -from the registered ``SeedDatasetProvider`` into memory. If a configured dataset -name still yields nothing, the resolver raises loudly rather than silently skipping it. -Inline configs (``seeds=`` / ``seed_groups=``) never touch memory. +Memory is the source of truth. ``prepare_async`` checks all sources before fetching +missing registered datasets. Read methods never fetch or write. Selection applies each +source's limit before the combined ``max_total`` limit. Inline configurations never +touch memory. """ from __future__ import annotations +import copy import random -from contextlib import contextmanager -from contextvars import ContextVar from dataclasses import dataclass from enum import Enum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast +from typing import TYPE_CHECKING, Any, Literal, Self, TypeVar, cast from pyrit.common import forward_init_parameters +from pyrit.common.deprecation import print_deprecation_message from pyrit.memory import CentralMemory from pyrit.models import ( AllAvailableDatasetSize, @@ -48,9 +47,10 @@ group_seeds_into_attack_groups, scenario_dataset_size_from_limit, ) +from pyrit.models.dataset_limit import DatasetLimit, ResolvedDatasetLimit, normalize_dataset_limit if TYPE_CHECKING: - from collections.abc import Callable, Iterator, Sequence + from collections.abc import Callable, Sequence from pyrit.memory import MemoryInterface @@ -61,17 +61,58 @@ # Internal helper TypeVar for size-capping any homogeneous list. _ItemT = TypeVar("_ItemT") -_AUTO_FETCH_ALLOWED: ContextVar[bool] = ContextVar("dataset_auto_fetch_allowed", default=True) -@contextmanager -def read_only_dataset_resolution() -> Iterator[None]: - """Disable dataset auto-fetch persistence within the current async context.""" - token = _AUTO_FETCH_ALLOWED.set(False) +class _Unset(Enum): + VALUE = "unset" + + +class DatasetFetchPolicy(str, Enum): + """When preparation may load a registered dataset into memory.""" + + NEVER = "never" + IF_MISSING = "if_missing" + + +def _resolve_limit(*, name: str, value: DatasetLimit | _Unset, default: ResolvedDatasetLimit) -> ResolvedDatasetLimit: try: - yield - finally: - _AUTO_FETCH_ALLOWED.reset(token) + limit = normalize_dataset_limit(None if isinstance(value, _Unset) else value) + except ValueError as exc: + raise DatasetConstraintError(f"'{name}': {exc}") from exc + return default if limit == "default" else limit + + +def _deprecated_argument(*, old: str, new: str) -> None: + print_deprecation_message( + old_item=f"DatasetConfiguration({old})", + new_item=f"DatasetConfiguration({new})", + removed_in="1.4.0", + ) + + +@dataclass(frozen=True, kw_only=True) +class DatasetSource: + """A named dataset, with optional overrides for its selection limit and fetch policy.""" + + name: str + max_size: DatasetLimit = "default" + fetch: DatasetFetchPolicy | None = None + + def __post_init__(self) -> None: + """ + Validate source options without reading memory or providers. + + Raises: + DatasetConstraintError: If the name, limit, or policy is invalid. + """ + if not isinstance(self.name, str) or not self.name.strip(): + raise DatasetConstraintError("Dataset source names must be non-empty strings.") + try: + object.__setattr__(self, "max_size", normalize_dataset_limit(self.max_size)) + except ValueError as exc: + raise DatasetConstraintError(f"'max_size': {exc}") from exc + if self.fetch is not None and not isinstance(self.fetch, DatasetFetchPolicy): + raise DatasetConstraintError("'fetch' must be a DatasetFetchPolicy.") class DatasetSourceKind(Enum): @@ -79,8 +120,8 @@ class DatasetSourceKind(Enum): How a ``DatasetConfiguration``'s seeds were sourced. Only two cases matter to validators: seeds supplied inline by the caller, versus - seeds loaded from memory by dataset name (auto-fetched into memory first when - missing). This lets a constraint require or forbid inline data -- e.g. a CLI + seeds loaded from memory by dataset name (prepared explicitly before reading). + This lets a constraint require or forbid inline data -- e.g. a CLI ``--objectives`` flag that must be passed inline rather than via a named dataset. """ @@ -97,7 +138,7 @@ class ResolvedDataset: supplied (inline vs named dataset), and which dataset names contributed. Args: - seeds (Sequence[Seed]): The resolved seeds (before ``max_dataset_size`` sampling). + seeds (Sequence[Seed]): The resolved seeds before sampling. source_kind (DatasetSourceKind): How the configuration was sourced. dataset_names (tuple[str, ...]): The configured dataset names that contributed seeds, in configuration order. Empty for inline ``seeds`` / ``seed_groups``. @@ -273,20 +314,14 @@ class DatasetConfiguration: """ Configuration describing where a scenario's seeds come from. - This base class handles resolution, fetching, validation, and sampling. + This base class separates preparation from read-only resolution and validation. ``DatasetAttackConfiguration`` is the concrete subclass most scenarios use; it groups the resolved seeds into ``AttackSeedGroup`` s. A configuration draws from exactly one source: - ``seeds`` -- an explicit, inline list of seeds (never touches memory). - ``seed_groups`` -- explicit, inline seed groups (never touches memory). - - ``dataset_names`` -- names looked up in memory; missing names are fetched from the - registered ``SeedDatasetProvider`` when ``auto_fetch`` is enabled. - - Resolution reads memory (the source of truth) and, per dataset name, fetches from the - provider when missing and ``auto_fetch`` is set. If a configured name still yields no - seeds, ``_collect_seeds_for_dataset_async`` raises ``DatasetConstraintError`` -- failures - are loud, not silently skipped. + - ``sources`` -- named datasets read from memory after explicit preparation. Constraints are expressed through a single mechanism -- ``validators`` -- so there is one place to look. Customize behavior through small seams without re-implementing @@ -303,11 +338,15 @@ def __init__( *, seeds: Sequence[Seed] | None = None, seed_groups: list[SeedGroup] | None = None, + sources: Sequence[DatasetSource] | None = None, + max_per_dataset: DatasetLimit = "default", + max_total: DatasetLimit | _Unset = _Unset.VALUE, + fetch: DatasetFetchPolicy | _Unset = _Unset.VALUE, dataset_names: list[str] | None = None, - max_dataset_size: int | None = None, + max_dataset_size: DatasetLimit | _Unset = _Unset.VALUE, filters: dict[str, list[str]] | None = None, validators: Sequence[Callable[[ResolvedDataset], None]] | None = None, - auto_fetch: bool = True, + auto_fetch: bool | _Unset = _Unset.VALUE, ) -> None: """ Initialize a DatasetConfiguration. @@ -316,43 +355,195 @@ def __init__( seeds (Sequence[Seed] | None): Explicit, inline seeds (never touches memory). seed_groups (list[SeedGroup] | None): Explicit, inline seed groups (never touches memory). - dataset_names (list[str] | None): Names of datasets to load from memory. - max_dataset_size (int | None): If set, randomly samples up to this many items - from the resolved dataset (without replacement). + sources (Sequence[DatasetSource] | None): Named datasets to prepare and read. + max_per_dataset (DatasetLimit): Source cap. None, empty, and "default" use the default; + "all" removes the cap. + max_total (DatasetLimit | _Unset): Combined cap, with the same default/all semantics. + fetch (DatasetFetchPolicy | _Unset): Default preparation policy; defaults to IF_MISSING. + dataset_names (list[str] | None): Deprecated alias for sources. + max_dataset_size (DatasetLimit | _Unset): Deprecated alias for max_total. filters (dict[str, list[str]] | None): Filters passed to ``MemoryInterface.get_seeds`` when resolving named datasets (e.g. ``{"harm_categories": ["cyber"]}``). - Applied before ``max_dataset_size`` sampling; ignored for inline seeds. + Applied before sampling; ignored for inline seeds. validators (Sequence[Callable[[ResolvedDataset], None]] | None): Constraint callbacks run against the resolved dataset; each raises on violation. These are appended to the subclass's ``_default_validators``. - auto_fetch (bool): When True (default), a configured dataset name that is not - in memory is fetched from the registered ``SeedDatasetProvider`` into - memory before resolving. Set False for strict "must already be in memory". + auto_fetch (bool | _Unset): Deprecated preparation-policy alias. Raises: - ValueError: If more than one of seeds/seed_groups/dataset_names is set. - ValueError: If max_dataset_size is less than 1. - """ - sources = [src for src in (seeds, seed_groups, dataset_names) if src is not None] - if len(sources) > 1: - raise ValueError( - "Only one of 'seeds', 'seed_groups', or 'dataset_names' can be set. " - "Use 'seeds'/'seed_groups' to provide inline data, or 'dataset_names' to load from memory." - ) - - if max_dataset_size is not None and max_dataset_size < 1: - raise ValueError("'max_dataset_size' must be a positive integer (>= 1).") - + ValueError: If source definitions or aliases conflict, or selection options are invalid. + """ + if sum(src is not None for src in (seeds, seed_groups, sources, dataset_names)) > 1: + raise ValueError("Only one of 'seeds', 'seed_groups', 'sources', or 'dataset_names' can be set.") + if dataset_names is not None: + _deprecated_argument(old="dataset_names=...", new="sources=...") + sources = [DatasetSource(name=name) for name in dataset_names] + if not isinstance(max_dataset_size, _Unset): + if not isinstance(max_total, _Unset): + raise ValueError("Use only one of 'max_dataset_size' and 'max_total'.") + _deprecated_argument(old="max_dataset_size=...", new="max_total=...") + max_total = max_dataset_size + if not isinstance(auto_fetch, _Unset): + if not isinstance(fetch, _Unset): + raise ValueError("Use only one of 'auto_fetch' and 'fetch'.") + if type(auto_fetch) is not bool: + raise ValueError("'auto_fetch' must be a bool.") + _deprecated_argument(old="auto_fetch=...", new="fetch=...") + fetch = DatasetFetchPolicy.IF_MISSING if auto_fetch else DatasetFetchPolicy.NEVER self._seeds = list(seeds) if seeds is not None else None self._seed_groups = list(seed_groups) if seed_groups is not None else None - self._dataset_names = list(dataset_names) if dataset_names is not None else None - self.max_dataset_size = max_dataset_size + self.sources = tuple(sources or ()) + self.max_per_dataset = max_per_dataset + self.max_total = None if isinstance(max_total, _Unset) else max_total + self.fetch = DatasetFetchPolicy.IF_MISSING if isinstance(fetch, _Unset) else fetch self._filters: dict[str, list[str]] = dict(filters or {}) self._validators: list[Callable[[ResolvedDataset], None]] = [ *self._default_validators(), *(list(validators) if validators else []), ] - self._auto_fetch = auto_fetch + self._validate_selection_options() + + def _validate_selection_options(self) -> None: + if not isinstance(self.fetch, DatasetFetchPolicy): + raise DatasetConstraintError("'fetch' must be a DatasetFetchPolicy.") + if (self._seeds is not None or self._seed_groups is not None) and self.max_per_dataset != "all": + raise DatasetConstraintError("Inline seeds do not support per-dataset limits; use max_total.") + names = [source.name for source in self.sources] + if len(names) != len(set(names)): + raise DatasetConstraintError("Duplicate dataset source names are not allowed.") + + @property + def max_total(self) -> ResolvedDatasetLimit: + """The resolved combined cap, or "all" for no total cap.""" + return self._max_total + + @max_total.setter + def max_total(self, value: DatasetLimit) -> None: + self._max_total = _resolve_limit(name="max_total", value=value, default=self._default_max_total()) + + @property + def max_per_dataset(self) -> ResolvedDatasetLimit: + """The resolved inherited source cap, or "all" for no source cap.""" + return self._max_per_dataset + + @max_per_dataset.setter + def max_per_dataset(self, value: DatasetLimit) -> None: + self._max_per_dataset = _resolve_limit( + name="max_per_dataset", value=value, default=self._default_max_per_dataset() + ) + + def _default_max_total(self) -> ResolvedDatasetLimit: + return "all" + + def _default_max_per_dataset(self) -> ResolvedDatasetLimit: + return "all" + + @property + def max_dataset_size(self) -> ResolvedDatasetLimit: + """The deprecated alias for the total selection limit.""" + _deprecated_argument(old="max_dataset_size", new="max_total") + return self.max_total + + @max_dataset_size.setter + def max_dataset_size(self, value: DatasetLimit) -> None: + _deprecated_argument(old="max_dataset_size", new="max_total") + self.max_total = value + + def with_overrides( + self, + *, + sources: Sequence[DatasetSource] | _Unset = _Unset.VALUE, + max_per_dataset: DatasetLimit | _Unset = _Unset.VALUE, + max_total: DatasetLimit | _Unset = _Unset.VALUE, + fetch: DatasetFetchPolicy | _Unset = _Unset.VALUE, + filters: dict[str, list[str]] | None = None, + ) -> Self: + """ + Copy this configuration, preserving its concrete class and custom state. + + Default-valued limit overrides keep the current settings. "all" removes only that cap. + + Returns: + Self: An independent selection configuration with shared live objects. + + Raises: + DatasetConstraintError: If named sources replace inline seeds or options are invalid. + """ + result = copy.copy(self) + result._filters = {key: list(values) for key, values in self._filters.items()} + result._validators = list(self._validators) + if not isinstance(sources, _Unset): + if self.source_kind is DatasetSourceKind.INLINE: + raise DatasetConstraintError("Cannot replace inline seeds with named sources through overrides.") + result.sources = tuple(sources) + result.max_per_dataset = _resolve_limit( + name="max_per_dataset", value=max_per_dataset, default=self.max_per_dataset + ) + result.max_total = _resolve_limit(name="max_total", value=max_total, default=self.max_total) + if not isinstance(fetch, _Unset): + result.fetch = fetch + if filters is not None: + result.update_filters(filters=filters) + result._validate_selection_options() + return result + + def source_limit(self, name: str) -> ResolvedDatasetLimit: + """Return the effective limit for a configured objective source.""" + source = next(source for source in self.sources if source.name == name) + return _resolve_limit(name="max_size", value=source.max_size, default=self.max_per_dataset) + + @property + def has_sampling_limits(self) -> bool: + """Whether this configuration can select a subset of its groups.""" + return self.max_total != "all" or any(self.source_limit(source.name) != "all" for source in self.sources) + + def _preparation_sources(self) -> list[tuple[str, DatasetFetchPolicy]]: + return [(source.name, source.fetch or self.fetch) for source in self.sources] + + async def prepare_async(self) -> None: + """ + Check all sources, then populate missing registered datasets in memory. + + Raises: + DatasetConstraintError: If a missing source cannot be fetched or provider data is invalid. + """ + self.validate_configuration() + if self.source_kind is DatasetSourceKind.INLINE: + return + policies: dict[str, DatasetFetchPolicy] = {} + for name, policy in self._preparation_sources(): + if name in policies and policies[name] is not policy: + raise DatasetConstraintError(f"Dataset '{name}' has conflicting fetch policies.") + policies[name] = policy + if not policies: + return + stored_names = set(await self._memory.get_seed_dataset_names_async()) + missing: list[str] = [] + for name, policy in policies.items(): + if name in stored_names: + continue + if policy is DatasetFetchPolicy.NEVER: + raise DatasetConstraintError(f"Dataset '{name}' is missing and fetch is 'never'. Import it first.") + missing.append(name) + if not missing: + return + from pyrit.datasets.seed_datasets.seed_dataset_provider import SeedDatasetProvider + + providers = await SeedDatasetProvider.get_providers_by_name_async(dataset_names=missing) + unavailable = [name for name in missing if name not in providers] + if unavailable: + raise DatasetConstraintError(f"Datasets {unavailable} are missing and have no provider. Import them first.") + for name in missing: + dataset = await providers[name].fetch_dataset_async() + if ( + not dataset.seeds + or dataset.dataset_name != name + or any(seed.dataset_name != name for seed in dataset.seeds) + ): + raise DatasetConstraintError( + f"Provider for '{name}' must return non-empty seeds with that dataset name." + ) + await self._memory.add_seed_datasets_to_memory_async(datasets=[dataset], added_by="DatasetConfiguration") def _default_validators(self) -> list[Callable[[ResolvedDataset], None]]: """ @@ -389,7 +580,7 @@ def dataset_names(self) -> list[str]: Returns: list[str]: The dataset names, or an empty list when using inline seeds/groups. """ - return list(self._dataset_names or []) + return [source.name for source in self.sources] @property def source_kind(self) -> DatasetSourceKind: @@ -418,7 +609,13 @@ def filters(self) -> dict[str, list[str]]: def get_size_budget(self) -> ScenarioDatasetSizeEstimate: """Return the configured selection budget without reading or sampling seeds.""" - return scenario_dataset_size_from_limit(self.max_dataset_size) + if not self.sources: + return scenario_dataset_size_from_limit(self.max_total) + limits = [self.source_limit(source.name) for source in self.sources] + if any(limit == "all" for limit in limits): + return scenario_dataset_size_from_limit(self.max_total) + total = sum(limit for limit in limits if isinstance(limit, int)) + return BoundedDatasetSize(value=min(total, self.max_total) if self.max_total != "all" else total) def validate_configuration(self) -> None: """ @@ -427,8 +624,7 @@ def validate_configuration(self) -> None: Raises: DatasetConstraintError: If the selection cap is not positive. """ - if self.max_dataset_size is not None and self.max_dataset_size < 1: - raise DatasetConstraintError("'max_dataset_size' must be a positive integer (>= 1).") + self._validate_selection_options() def size_caps_by_dataset(self) -> dict[str, list[tuple[str, int, Literal["dataset", "configuration", "compound"]]]]: """ @@ -438,14 +634,16 @@ def size_caps_by_dataset(self) -> dict[str, list[tuple[str, int, Literal["datase dict[str, list[tuple[str, int, Literal]]]: Source name to ordered ``(cap label, count, provenance)`` entries. """ - if self.max_dataset_size is None: - return {} - names = self.dataset_names or [INLINE_DATASET_NAME] - if len(names) == 1: - cap = ("per-dataset cap", self.max_dataset_size, "dataset") - else: - cap = ("combined configuration cap", self.max_dataset_size, "configuration") - return {name: [cap] for name in names} + caps: dict[str, list[tuple[str, int, Literal["dataset", "configuration", "compound"]]]] = {} + for name in self.dataset_names or [INLINE_DATASET_NAME]: + entries: list[tuple[str, int, Literal["dataset", "configuration", "compound"]]] = [] + limit = self.source_limit(name) if self.sources else "all" + if limit != "all": + entries.append(("per-dataset cap", limit, "dataset")) + if self.max_total != "all": + entries.append(("combined configuration cap", self.max_total, "configuration")) + caps[name] = entries + return caps @property def _get_seeds_filters(self) -> dict[str, Any]: @@ -470,7 +668,7 @@ def update_filters(self, *, filters: dict[str, list[str]]) -> None: Args: filters (dict[str, list[str]]): Filters to merge, keyed by ``get_seeds`` kwarg name. """ - self._filters = {**self._filters, **filters} + self._filters = {**self._filters, **{key: list(values) for key, values in filters.items()}} # ========================================================================= # Resolution helpers @@ -480,8 +678,8 @@ async def _collect_named_seeds_async(self) -> dict[str, list[Seed]]: """ Collect seeds for each configured dataset name, keyed by name. - Each name is read from memory and -- when empty and ``auto_fetch`` is set -- fetched - from the provider; a name that still yields nothing raises loudly. + Each name is read from memory. A missing dataset requires explicit preparation + or import; a read never calls a provider. Returns: dict[str, list[Seed]]: Dataset name -> seeds, in configuration order (every value @@ -491,13 +689,13 @@ async def _collect_named_seeds_async(self) -> dict[str, list[Seed]]: DatasetConstraintError: If any configured dataset yields no seeds. """ result: dict[str, list[Seed]] = {} - for name in self._dataset_names or []: + for name in self.dataset_names: result[name] = await self._collect_seeds_for_dataset_async(dataset_name=name) return result async def _collect_seeds_for_dataset_async(self, *, dataset_name: str) -> list[Seed]: """ - Collect seeds for a single dataset name, fetching from the provider if needed. + Read seeds for a single dataset name without fetching or writing. Args: dataset_name (str): The dataset name to load. @@ -506,58 +704,21 @@ async def _collect_seeds_for_dataset_async(self, *, dataset_name: str) -> list[S list[Seed]: The seeds for ``dataset_name``. Raises: - DatasetConstraintError: If the dataset yields no seeds even after auto-fetch, or - if auto-fetch itself fails (the provider error is chained as the cause). + DatasetConstraintError: If the dataset is absent or its filters match no seeds. """ found = list(await self._memory.get_seeds_async(dataset_name=dataset_name, **self._get_seeds_filters)) - auto_fetch_allowed = self._auto_fetch and _AUTO_FETCH_ALLOWED.get() - if not found and auto_fetch_allowed: - try: - await self._fetch_dataset_async(dataset_name=dataset_name) - except Exception as exc: - raise DatasetConstraintError( - f"Dataset '{dataset_name}' could not be loaded: auto-fetch from the registered provider failed." - ) from exc - found = list(await self._memory.get_seeds_async(dataset_name=dataset_name, **self._get_seeds_filters)) if not found: unfiltered = await self._memory.get_seeds_async(dataset_name=dataset_name) if self._filters else [] if unfiltered: raise DatasetConstraintError( f"Dataset '{dataset_name}' has seeds, but none match the configured filters {self._filters}." ) - if auto_fetch_allowed: - hint = "auto-fetch from the registered provider did not populate it" - elif self._auto_fetch: - hint = "auto_fetch is disabled for read-only resolution" - else: - hint = "auto_fetch is disabled" raise DatasetConstraintError( - f"Dataset '{dataset_name}' could not be loaded: no seeds found in memory and {hint}." + f"Dataset '{dataset_name}' could not be loaded: no seeds found in memory. " + "Call prepare_async() before reading a new selection, or import the dataset." ) return found - async def _fetch_dataset_async(self, *, dataset_name: str) -> None: - """ - Populate memory from the registered provider for a single dataset (private). - - An unregistered name populates nothing and falls through to the caller's loud - empty-result handling. Provider errors (enumeration or fetch) propagate so the - caller can surface the root cause. Never samples or validates -- it only adds to - memory. - - Args: - dataset_name (str): The dataset name to fetch. - """ - # Local import to avoid an import cycle at package init time. - from pyrit.datasets.seed_datasets.seed_dataset_provider import SeedDatasetProvider - - registered = set(await SeedDatasetProvider.get_all_dataset_names_async()) - if dataset_name not in registered: - return - - datasets = await SeedDatasetProvider.fetch_datasets_async(dataset_names=[dataset_name]) - await self._memory.add_seed_datasets_to_memory_async(datasets=datasets, added_by="DatasetConfiguration") - def validate(self, resolved: ResolvedDataset) -> None: """ Validate the resolved dataset against every configured validator. @@ -577,18 +738,38 @@ def validate(self, resolved: ResolvedDataset) -> None: def _apply_max_dataset_size(self, items: list[_ItemT]) -> list[_ItemT]: """ - Apply ``max_dataset_size`` sampling without replacement. + Apply ``max_total`` sampling without replacement. Args: items (list[_ItemT]): The items to potentially sample from. Returns: list[_ItemT]: The original list, or a random sample of up to - ``max_dataset_size`` unique items. + ``max_total`` unique items. """ - if self.max_dataset_size is None or len(items) <= self.max_dataset_size: + if self.max_total == "all" or len(items) <= self.max_total: return items - return random.sample(items, self.max_dataset_size) + return random.sample(items, self.max_total) + + +@dataclass(frozen=True) +class _ResolvedAttackGroups: + """A per-read snapshot that validates compound populations before any sampling.""" + + configuration: DatasetAttackConfiguration + groups: dict[str, list[AttackSeedGroup]] + children: tuple[_ResolvedAttackGroups, ...] = () + + def select(self, *, apply_sampling: bool) -> dict[str, list[AttackSeedGroup]]: + if not apply_sampling: + return self.groups + groups = self.groups + if self.children: + groups = {} + for child in self.children: + for name, population in child.select(apply_sampling=True).items(): + groups.setdefault(name, []).extend(population) + return self.configuration._sample_groups_by_dataset(groups) class DatasetAttackConfiguration(DatasetConfiguration): @@ -596,19 +777,14 @@ class DatasetAttackConfiguration(DatasetConfiguration): A ``DatasetConfiguration`` that groups resolved seeds into attack groups. This is the default most scenarios use: scenarios run over ``AttackSeedGroup`` s - (each carrying exactly one objective plus optional prompts). ``max_dataset_size`` is a - single global budget for the configuration; both resolvers apply it the same way: + (each carrying exactly one objective plus optional prompts). Both resolvers apply + source limits followed by one ``max_total`` limit: - ``get_attack_seed_groups_async`` -- a flat ``list[AttackSeedGroup]``, sampled globally over all built groups. - ``get_attack_groups_by_dataset_async`` -- the same globally sampled groups, keyed by dataset name, used when a scenario fans atomic attacks out per (technique, dataset). - To draw an independent budget from *each* of several datasets (the old "N per dataset" - behavior), compose one child per dataset with ``CompoundDatasetAttackConfiguration`` -- - e.g. ``CompoundDatasetAttackConfiguration.per_dataset(dataset_names=[...], max_dataset_size=4)`` - -- rather than relying on a single config to special-case dataset names. - Both run ``validators`` against the full resolved seed set before sampling. Override ``_build_attack_groups`` to change how raw seeds become attack groups @@ -617,16 +793,51 @@ class DatasetAttackConfiguration(DatasetConfiguration): """ @forward_init_parameters - def __init__(self, *, max_dataset_size: int | None = 5, **kwargs: Any) -> None: + def __init__( + self, + *, + sampling_scope: Literal["per_dataset", "total_only"] = "per_dataset", + **kwargs: Any, + ) -> None: """ Configure scenario attack groups with a finite default selection cap. Args: - max_dataset_size (int | None): Maximum selected attack groups. Defaults to 5; - explicit None retains the full population. - **kwargs (Any): Dataset source, filters, and validation options. + sampling_scope (Literal["per_dataset", "total_only"]): Apply per-dataset caps before the total cap, + or cap only the combined attack groups. Use total_only for complete ingredient populations. + **kwargs (Any): Dataset source, limits, filters, and validation options. Default limits are + 5 per named objective source, or 5 total for inline groups. Use "all" to remove a limit. """ - super().__init__(max_dataset_size=max_dataset_size, **kwargs) + self.sampling_scope = sampling_scope + super().__init__(**kwargs) + + def _default_max_total(self) -> ResolvedDatasetLimit: + return 5 if self._seeds is not None or self._seed_groups is not None else "all" + + def _default_max_per_dataset(self) -> ResolvedDatasetLimit: + inline = self._seeds is not None or self._seed_groups is not None + return 5 if self.sampling_scope == "per_dataset" and not inline else "all" + + def _validate_selection_options(self) -> None: + """ + Validate limits against the configured sampling scope without reading datasets. + + Raises: + DatasetConstraintError: If the sampling scope is invalid or conflicts with per-dataset limits. + """ + super()._validate_selection_options() + if self.sampling_scope not in ("per_dataset", "total_only"): + raise DatasetConstraintError("'sampling_scope' must be 'per_dataset' or 'total_only'.") + if self.sampling_scope == "total_only" and ( + self.max_per_dataset != "all" or any(self.source_limit(source.name) != "all" for source in self.sources) + ): + raise DatasetConstraintError("Total-only sampling does not support per-dataset limits; use max_total.") + + def get_size_budget(self) -> ScenarioDatasetSizeEstimate: + """Return the selection budget, excluding ingredient row counts.""" + if self.sampling_scope == "total_only": + return scenario_dataset_size_from_limit(self.max_total) + return super().get_size_budget() def _build_attack_groups(self, seeds: list[Seed]) -> list[AttackSeedGroup]: """ @@ -666,7 +877,7 @@ async def _build_groups_by_dataset_async(self) -> tuple[dict[str, list[AttackSee Inline configs preserve their explicit grouping under the ``INLINE_DATASET_NAME`` label (they are not flattened and regrouped). Named datasets reuse ``_collect_named_seeds_async`` - (auto-fetch + loud empty handling) and run each dataset's seeds through + (read-only collection) and run each dataset's seeds through ``_build_attack_groups``. Returns: @@ -696,12 +907,11 @@ async def get_attack_seed_groups_async(self, *, apply_sampling: bool = True) -> """ Resolve the configured dataset into a flat ``list[AttackSeedGroup]``. - Builds attack groups (inline or from memory, auto-fetching missing datasets), - validates the full resolved seed set, then samples ``max_dataset_size`` globally - over all built groups. + Builds attack groups from inline data or memory, validates the full resolved + seed set, then samples each source and the combined population. Args: - apply_sampling (bool): When True (default), apply ``max_dataset_size`` sampling. + apply_sampling (bool): When True (default), apply source and total sampling. Pass False to resolve the full, deterministic dataset with no ``random.sample`` draw -- used on resume so the persisted objective subset can be reconstructed exactly rather than intersected against a fresh (divergent) sample. @@ -714,16 +924,8 @@ async def get_attack_seed_groups_async(self, *, apply_sampling: bool = True) -> DatasetConstraintError: If a configured dataset yields no seeds, the resolved dataset fails validation, or no attack groups could be built. """ - self.validate_configuration() - groups_by_dataset, resolved = await self._build_groups_by_dataset_async() - self.validate(resolved) - groups = [group for groups in groups_by_dataset.values() for group in groups] - if apply_sampling: - groups = self._apply_max_dataset_size(groups) - if not groups: - names = ", ".join(self._dataset_names) if self._dataset_names else "" - raise DatasetConstraintError(f"Resolved attack-group dataset is empty (datasets: {names}).") - return groups + grouped = await self.get_attack_groups_by_dataset_async(apply_sampling=apply_sampling) + return [group for groups in grouped.values() for group in groups] async def get_attack_groups_by_dataset_async( self, *, apply_sampling: bool = True @@ -731,14 +933,12 @@ async def get_attack_groups_by_dataset_async( """ Resolve attack groups keyed by dataset name, globally sampled. - Inline configs resolve under the ``INLINE_DATASET_NAME`` label. Builds attack groups - (auto-fetching missing datasets), validates the full resolved seed set, then applies - ``max_dataset_size`` as one global budget across all datasets -- the survivors stay - keyed by their originating dataset. For an independent budget per dataset, compose - ``CompoundDatasetAttackConfiguration.per_dataset(...)`` instead. + Inline configs resolve under the ``INLINE_DATASET_NAME`` label. Validate the + full seed set, apply source limits, then apply ``max_total`` to the union. + Survivors retain their source association. Args: - apply_sampling (bool): When True (default), apply ``max_dataset_size`` sampling. + apply_sampling (bool): When True (default), apply source and total sampling. Pass False to resolve the full, deterministic dataset with no ``random.sample`` draw -- used on resume so the persisted objective subset can be reconstructed exactly rather than intersected against a fresh (divergent) sample. @@ -751,23 +951,27 @@ async def get_attack_groups_by_dataset_async( DatasetConstraintError: If a configured dataset yields no seeds, the resolved dataset fails validation, or no attack groups could be built. """ - self.validate_configuration() - groups_by_dataset, resolved = await self._build_groups_by_dataset_async() - self.validate(resolved) - sampled = self._sample_groups_by_dataset(groups_by_dataset) if apply_sampling else groups_by_dataset + selection = await self._resolve_attack_groups_async() + sampled = selection.select(apply_sampling=apply_sampling) result = {name: groups for name, groups in sampled.items() if groups} if not result: - names = ", ".join(self._dataset_names) if self._dataset_names else "" + names = ", ".join(self.dataset_names) if self.dataset_names else "" raise DatasetConstraintError(f"Resolved attack-group dataset is empty (datasets: {names}).") return result + async def _resolve_attack_groups_async(self) -> _ResolvedAttackGroups: + self.validate_configuration() + groups, resolved = await self._build_groups_by_dataset_async() + self.validate(resolved) + return _ResolvedAttackGroups(configuration=self, groups=groups) + def _sample_groups_by_dataset( self, groups_by_dataset: dict[str, list[AttackSeedGroup]] ) -> dict[str, list[AttackSeedGroup]]: """ - Apply ``max_dataset_size`` as one global budget across datasets, preserving keys. + Apply source caps, then the total cap, preserving dataset keys. - Flattens every ``(dataset_name, group)`` pair, samples up to ``max_dataset_size`` + Flattens every ``(dataset_name, group)`` pair, samples up to ``max_total`` across the union, then regroups the survivors under their originating dataset name. Args: @@ -776,12 +980,21 @@ def _sample_groups_by_dataset( Returns: dict[str, list[AttackSeedGroup]]: The globally sampled groups, still keyed by dataset. """ - pairs = [(name, group) for name, groups in groups_by_dataset.items() for group in groups] + limited = groups_by_dataset + if self.sampling_scope == "per_dataset" and self.sources: + limited = { + name: self._sample_source_groups(name=name, groups=groups) for name, groups in groups_by_dataset.items() + } + pairs = [(name, group) for name, groups in limited.items() for group in groups] result: dict[str, list[AttackSeedGroup]] = {} for name, group in self._apply_max_dataset_size(pairs): result.setdefault(name, []).append(group) return result + def _sample_source_groups(self, *, name: str, groups: list[AttackSeedGroup]) -> list[AttackSeedGroup]: + limit = self.source_limit(name) + return random.sample(groups, limit) if limit != "all" and len(groups) > limit else groups + class CompoundDatasetAttackConfiguration(DatasetAttackConfiguration): """ @@ -792,17 +1005,16 @@ class CompoundDatasetAttackConfiguration(DatasetAttackConfiguration): budgets or shaping -- for example "up to 4 attack groups from *each* of several datasets" (see ``per_dataset``), or pairing one dataset's objectives with another dataset's prompts. - A single ``DatasetAttackConfiguration`` applies ``max_dataset_size`` as one global budget; - per-dataset budgets are expressed by composing one child per dataset rather than by - special-casing dataset names inside a single configuration. An optional compound-level - ``max_dataset_size`` caps the combined result on top of each child's own sampling. + Use a compound for independent shaping or validators. Simple per-dataset budgets + use ``DatasetSource`` instead. ``max_total`` caps the combined child selections. """ def __init__( self, *, configurations: Sequence[DatasetAttackConfiguration], - max_dataset_size: int | None = None, + max_total: DatasetLimit | _Unset = _Unset.VALUE, + max_dataset_size: DatasetLimit | _Unset = _Unset.VALUE, validators: Sequence[Callable[[ResolvedDataset], None]] | None = None, ) -> None: """ @@ -811,8 +1023,8 @@ def __init__( Args: configurations (Sequence[DatasetAttackConfiguration]): The child configurations to combine; each resolves and samples independently. Must be non-empty. - max_dataset_size (int | None): Optional cap applied to the *combined* result, on - top of each child's own sampling. + max_total (DatasetLimit | _Unset): Optional cap on the combined child selections. + max_dataset_size (DatasetLimit | _Unset): Deprecated alias for max_total. validators (Sequence[Callable[[ResolvedDataset], None]] | None): Validators run against the combined resolved seeds, in addition to each child's validators. @@ -821,15 +1033,54 @@ def __init__( """ if not configurations: raise ValueError("CompoundDatasetAttackConfiguration requires at least one child configuration.") - super().__init__(max_dataset_size=max_dataset_size, validators=validators) + super().__init__( + max_total=max_total, max_dataset_size=max_dataset_size, max_per_dataset="all", validators=validators + ) self._configurations = list(configurations) + @property + def has_sampling_limits(self) -> bool: + """Whether this compound or any child can select a subset.""" + return super().has_sampling_limits or any(child.has_sampling_limits for child in self._configurations) + + def _preparation_sources(self) -> list[tuple[str, DatasetFetchPolicy]]: + return [source for child in self._configurations for source in child._preparation_sources()] + + def with_overrides( + self, + *, + sources: Sequence[DatasetSource] | _Unset = _Unset.VALUE, + max_per_dataset: DatasetLimit | _Unset = _Unset.VALUE, + max_total: DatasetLimit | _Unset = _Unset.VALUE, + fetch: DatasetFetchPolicy | _Unset = _Unset.VALUE, + filters: dict[str, list[str]] | None = None, + ) -> Self: + """ + Copy a compound and its children without rebuilding their custom classes. + + Returns: + Self: The copied compound and child selection settings. + + Raises: + DatasetConstraintError: If source replacement is requested on the compound. + """ + if not isinstance(sources, _Unset): + raise DatasetConstraintError("Replace named sources on individual compound children, not the compound.") + result = super().with_overrides(max_total=max_total) + result._configurations = [ + child.with_overrides(max_per_dataset=max_per_dataset, fetch=fetch, filters=filters) + for child in self._configurations + ] + if filters is not None: + result._filters.update({key: list(values) for key, values in filters.items()}) + return result + @classmethod def per_dataset( cls, *, dataset_names: Sequence[str], - max_dataset_size: int | None = 5, + max_dataset_size: DatasetLimit = "default", auto_fetch: bool = True, filters: dict[str, list[str]] | None = None, validators: Sequence[Callable[[ResolvedDataset], None]] | None = None, @@ -842,8 +1093,8 @@ def per_dataset( Args: dataset_names (Sequence[str]): The dataset names; one child is built per name. - max_dataset_size (int | None): Per-dataset cap applied to each child. - Defaults to 5; pass None for unlimited children. + max_dataset_size (DatasetLimit): Per-dataset cap applied to each child. + Defaults to 5; pass "all" for unlimited children. auto_fetch (bool): Passed to each child (fetch missing datasets into memory). filters (dict[str, list[str]] | None): ``get_seeds`` filters applied to each child. validators (Sequence[Callable[[ResolvedDataset], None]] | None): Applied to each child. @@ -856,12 +1107,13 @@ def per_dataset( """ if not dataset_names: raise ValueError("per_dataset requires at least one dataset name.") + _deprecated_argument(old="CompoundDatasetAttackConfiguration.per_dataset(...)", new="sources=...") return cls( configurations=[ DatasetAttackConfiguration( - dataset_names=[name], - max_dataset_size=max_dataset_size, - auto_fetch=auto_fetch, + sources=[DatasetSource(name=name)], + max_per_dataset=max_dataset_size, + fetch=DatasetFetchPolicy.IF_MISSING if auto_fetch else DatasetFetchPolicy.NEVER, filters=filters, validators=validators, ) @@ -908,13 +1160,11 @@ def get_size_budget(self) -> ScenarioDatasetSizeEstimate: if isinstance(budget, IndeterminateDatasetSize): return budget if not all(isinstance(budget, BoundedDatasetSize) for budget in budgets): - if self.max_dataset_size is not None: - return BoundedDatasetSize(value=self.max_dataset_size) + if self.max_total != "all": + return BoundedDatasetSize(value=self.max_total) return AllAvailableDatasetSize() total = sum(budget.value for budget in budgets if isinstance(budget, BoundedDatasetSize)) - return BoundedDatasetSize( - value=min(total, self.max_dataset_size) if self.max_dataset_size is not None else total - ) + return BoundedDatasetSize(value=min(total, self.max_total) if self.max_total != "all" else total) def validate_configuration(self) -> None: """Check the compound and every child without resolving dataset contents.""" @@ -933,9 +1183,9 @@ def size_caps_by_dataset(self) -> dict[str, list[tuple[str, int, Literal["datase for child in self._configurations: for name, child_caps in child.size_caps_by_dataset().items(): caps.setdefault(name, []).extend(child_caps) - if self.max_dataset_size is not None: + if self.max_total != "all": for name in self.dataset_names or [INLINE_DATASET_NAME]: - caps.setdefault(name, []).append(("combined compound cap", self.max_dataset_size, "compound")) + caps.setdefault(name, []).append(("combined compound cap", self.max_total, "compound")) return caps def update_filters(self, *, filters: dict[str, list[str]]) -> None: @@ -971,12 +1221,8 @@ async def get_attack_seed_groups_async(self, *, apply_sampling: bool = True) -> Raises: DatasetConstraintError: If a child yields nothing, or the combined result fails validation. """ - self.validate_configuration() - groups: list[AttackSeedGroup] = [] - for child in self._configurations: - groups.extend(await child.get_attack_seed_groups_async(apply_sampling=apply_sampling)) - self.validate(self._resolved_from_groups(groups)) - return self._apply_max_dataset_size(groups) if apply_sampling else groups + grouped = await self.get_attack_groups_by_dataset_async(apply_sampling=apply_sampling) + return [group for groups in grouped.values() for group in groups] async def get_attack_groups_by_dataset_async( self, *, apply_sampling: bool = True @@ -995,14 +1241,17 @@ async def get_attack_groups_by_dataset_async( Raises: DatasetConstraintError: If a child yields nothing, or the combined result fails validation. """ + return await super().get_attack_groups_by_dataset_async(apply_sampling=apply_sampling) + + async def _resolve_attack_groups_async(self) -> _ResolvedAttackGroups: self.validate_configuration() + children = tuple([await child._resolve_attack_groups_async() for child in self._configurations]) merged: dict[str, list[AttackSeedGroup]] = {} - for child in self._configurations: - child_groups = await child.get_attack_groups_by_dataset_async(apply_sampling=apply_sampling) - for name, groups in child_groups.items(): + for child in children: + for name, groups in child.groups.items(): merged.setdefault(name, []).extend(groups) self.validate(self._resolved_from_groups([group for groups in merged.values() for group in groups])) - return self._sample_groups_by_dataset(merged) if apply_sampling else merged + return _ResolvedAttackGroups(configuration=self, groups=merged, children=children) def _resolved_from_groups(self, groups: list[AttackSeedGroup]) -> ResolvedDataset: """ diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 6e8179369b..c97267ed6a 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -734,10 +734,8 @@ def _get_dataset_limit_input(self) -> DatasetLimitInput: """Return the editable limit, never an aggregate population budget.""" if not self.USES_DATASET_SIZE_LIMIT: return DatasetLimitInput(state=DatasetLimitState.NotApplicable) - limit = self._dataset_config.max_dataset_size - return ( - DatasetLimitInput(state=DatasetLimitState.Value, value=limit) if limit is not None else DatasetLimitInput() - ) + limit = self._dataset_config.max_total + return DatasetLimitInput(state=DatasetLimitState.Value, value=limit) if limit != "all" else DatasetLimitInput() def _get_run_size_budget(self) -> ScenarioDatasetSizeEstimate: """Return the scenario's configured population budget.""" @@ -813,7 +811,7 @@ def _resolve_runtime_configuration(self, *, require_objective_target: bool) -> N dataset_config = params.get("dataset_config") self._dataset_config_provided = dataset_config is not None - self._dataset_config = dataset_config if dataset_config else self._default_dataset_config + self._dataset_config = (dataset_config if dataset_config else self._default_dataset_config).with_overrides() self._max_concurrency = params.get("max_concurrency", 4) self._max_retries = params.get("max_retries", 0) self._memory_labels = params.get("memory_labels") or {} @@ -837,8 +835,17 @@ def _resolve_runtime_configuration(self, *, require_objective_target: bool) -> N self._validate_runtime_configuration() def _validate_runtime_configuration(self) -> None: - """Check resolved parameters for preview and launch without reading datasets.""" + """ + Check resolved parameters for preview and launch without reading datasets. + + Raises: + ValueError: If an ingredient-only scenario receives per-dataset limits. + """ self._dataset_config.validate_configuration() + if not self.USES_DATASET_SIZE_LIMIT and any( + self._dataset_config.source_limit(source.name) != "all" for source in self._dataset_config.sources + ): + raise ValueError("This scenario uses ingredient datasets and does not support per-dataset limits.") @final async def initialize_async(self) -> None: @@ -859,8 +866,8 @@ async def initialize_async(self) -> None: If a scenario_result_id was provided in __init__, this method will check if it exists in memory and validate that the stored scenario matches the current configuration. - If it matches, the scenario will resume from prior progress. If it doesn't match or - doesn't exist, a new scenario result will be created. + If it matches, the scenario will resume from prior progress without preparation. + If it does not match or does not exist, initialization raises an error. The common run inputs read from the bag are ``objective_target`` (a ``PromptTarget`` instance or a registered target name resolved against ``TargetRegistry``), @@ -889,6 +896,18 @@ async def initialize_async(self) -> None: # replayed by _apply_persisted_objectives. Re-drawing a fresh random.sample here would # diverge from the persisted hashes and abort resume whenever max_dataset_size is set. is_resume = self._scenario_result_id is not None + existing_results = [] + if is_resume: + existing_results = await self._memory.get_scenario_results_async( + scenario_result_ids=[self._scenario_result_id] + ) + if not existing_results: + raise ValueError( + f"Scenario result id '{self._scenario_result_id}' not found in memory. " + "Drop scenario_result_id to start a new scenario." + ) + else: + await self._dataset_config.prepare_async() seed_groups_by_dataset = await self._resolve_seed_groups_by_dataset_async(apply_sampling=not is_resume) context = self._build_scenario_context(seed_groups_by_dataset=seed_groups_by_dataset) self._atomic_attacks = await self._build_atomic_attacks_async(context=context) @@ -902,16 +921,6 @@ async def initialize_async(self) -> None: # rather than a silent restart, so the original progress isn't orphaned without # the user knowing. if self._scenario_result_id: - existing_results = await self._memory.get_scenario_results_async( - scenario_result_ids=[self._scenario_result_id] - ) - - if not existing_results: - raise ValueError( - f"Scenario result id '{self._scenario_result_id}' not found in memory. " - f"Drop scenario_result_id to start a new scenario." - ) - self._validate_stored_scenario( stored_result=existing_results[0], current_identifier=scenario_identifier, @@ -962,20 +971,20 @@ def _build_initial_scenario_metadata(self) -> dict[str, Any]: """ Build the metadata dict persisted with a freshly-created ``ScenarioResult``. - When ``max_dataset_size`` is in effect, the dataset config draws an + When a source or total limit is in effect, the dataset config draws an unseeded ``random.sample`` and the chosen subset would silently change on the next run (e.g. a resume). To make resume reliable, snapshot the chosen objective hashes here so the next ``_setup_scenario_async`` can replay them via ``keep_seed_groups_with_hashes``. - The normalized run plan is always stored. When ``max_dataset_size`` is not - set, only the run plan is needed because the full dataset is deterministic. + The normalized run plan is always stored. When no sampling limits are set, + only the run plan is needed because the full dataset is deterministic. Returns: dict[str, Any]: Metadata payload for the new ScenarioResult. """ metadata: dict[str, Any] = {} - if getattr(self._dataset_config, "max_dataset_size", None) is not None: + if self._dataset_config.has_sampling_limits: hashes: list[str] = [] seen: set[str] = set() for aa in self._atomic_attacks: diff --git a/pyrit/scenario/scenarios/adaptive/text_adaptive.py b/pyrit/scenario/scenarios/adaptive/text_adaptive.py index ad1c468312..de4df73041 100644 --- a/pyrit/scenario/scenarios/adaptive/text_adaptive.py +++ b/pyrit/scenario/scenarios/adaptive/text_adaptive.py @@ -19,7 +19,7 @@ from pyrit.common import apply_defaults from pyrit.models.parameter import Parameter from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry -from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration, DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.scenarios.adaptive.adaptive_scenario import AdaptiveScenario if TYPE_CHECKING: @@ -113,7 +113,9 @@ def required_datasets(cls) -> list[str]: @classmethod def default_dataset_config(cls) -> DatasetAttackConfiguration: """Return the default dataset config (required datasets, capped at 4 per dataset).""" - return CompoundDatasetAttackConfiguration.per_dataset(dataset_names=cls.required_datasets(), max_dataset_size=4) + return DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in cls.required_datasets()], max_per_dataset=4 + ) @classmethod def additional_parameters(cls) -> list[Parameter]: diff --git a/pyrit/scenario/scenarios/airt/cyber.py b/pyrit/scenario/scenarios/airt/cyber.py index 2f622c7b41..a25d1f3950 100644 --- a/pyrit/scenario/scenarios/airt/cyber.py +++ b/pyrit/scenario/scenarios/airt/cyber.py @@ -9,7 +9,7 @@ from pyrit.common import apply_defaults from pyrit.common.path import SCORER_SEED_PROMPT_PATH -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_matrix_atomic_attacks from pyrit.scenario.core.scenario import Scenario @@ -106,7 +106,9 @@ def __init__( version=self.VERSION, objective_scorer=self._objective_scorer, technique_class=technique_class, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["airt_malware"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["airt_malware"]], max_per_dataset="all", max_total=4 + ), scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/airt/jailbreak.py b/pyrit/scenario/scenarios/airt/jailbreak.py index 9f8e1472d8..7d1db6847a 100644 --- a/pyrit/scenario/scenarios/airt/jailbreak.py +++ b/pyrit/scenario/scenarios/airt/jailbreak.py @@ -22,7 +22,7 @@ from pyrit.prompt_target import CapabilityName from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import ( MatrixAtomicAttackBuilder, build_baseline_atomic_attack, @@ -243,7 +243,9 @@ def __init__( super().__init__( version=self.VERSION, technique_class=technique_class, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["harmbench"]], max_per_dataset="all", max_total=4 + ), objective_scorer=self._objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/airt/leakage.py b/pyrit/scenario/scenarios/airt/leakage.py index 6df8c914a8..f8c4d89c64 100644 --- a/pyrit/scenario/scenarios/airt/leakage.py +++ b/pyrit/scenario/scenarios/airt/leakage.py @@ -10,7 +10,7 @@ from pyrit.common import apply_defaults from pyrit.common.path import SCORER_SEED_PROMPT_PATH from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_matrix_atomic_attacks from pyrit.scenario.core.scenario import Scenario @@ -118,7 +118,9 @@ def __init__( super().__init__( version=self.VERSION, technique_class=technique_class, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["airt_leakage"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["airt_leakage"]], max_per_dataset="all", max_total=4 + ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/airt/multilingual.py b/pyrit/scenario/scenarios/airt/multilingual.py index b48ffbe124..7c316a5df2 100644 --- a/pyrit/scenario/scenarios/airt/multilingual.py +++ b/pyrit/scenario/scenarios/airt/multilingual.py @@ -23,6 +23,7 @@ ScenarioTechnique, get_default_adversarial_target, ) +from pyrit.scenario.core.dataset_configuration import DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import ( MatrixAtomicAttackBuilder, build_baseline_atomic_attack, @@ -219,7 +220,9 @@ def __init__( version=self.VERSION, uses_default_adversarial_target=adversarial_chat is None, technique_class=technique_class, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=5), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["harmbench"]], max_per_dataset="all", max_total=5 + ), objective_scorer=self._objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/airt/psychosocial.py b/pyrit/scenario/scenarios/airt/psychosocial.py index 37e68e6781..a6db8359c9 100644 --- a/pyrit/scenario/scenarios/airt/psychosocial.py +++ b/pyrit/scenario/scenarios/airt/psychosocial.py @@ -32,11 +32,8 @@ CrescendoAttack, ) from pyrit.models import ( - AllAvailableDatasetSize, BoundedDatasetSize, - ScenarioDatasetSizeCap, ScenarioDatasetSizeEstimate, - ScenarioDatasetSummary, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, SeedPrompt, @@ -47,8 +44,9 @@ from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory from pyrit.scenario.core.dataset_configuration import ( - CompoundDatasetAttackConfiguration, DatasetAttackConfiguration, + DatasetConstraintError, + DatasetSource, ) from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack from pyrit.scenario.core.scenario import Scenario @@ -64,7 +62,6 @@ if TYPE_CHECKING: from collections.abc import Callable - from pyrit.models import AttackSeedGroup from pyrit.prompt_target import PromptTarget from pyrit.scenario.core.scenario_context import ScenarioContext from pyrit.score import FloatScaleScorer, TrueFalseScorer @@ -363,12 +360,11 @@ class Psychosocial(Scenario): (``all`` only) swaps the simulated base for a real ``CrescendoAttack``. One baseline per sub-harm is emitted (toggle with ``include_baseline``). - Dataset selection is bound to the sub-harms: the ``dataset_config`` parameter still tunes - ``max_dataset_size`` and sampling, but the dataset names are always the selected sub-harms' - datasets (``--dataset-names`` is ignored). + Dataset selection is bound to the selected sub-harms. Source caps apply independently, + then ``max_total`` caps the combined population. Unrelated dataset names are rejected. """ - VERSION: int = 4 + VERSION: int = 5 @classmethod def additional_parameters(cls) -> list[Parameter]: @@ -447,7 +443,7 @@ def __init__( uses_default_adversarial_target=adversarial_chat is None, technique_class=PsychosocialTechnique, default_dataset_config=DatasetAttackConfiguration( - dataset_names=[harm.dataset_name for harm in _SUB_HARMS], + sources=[DatasetSource(name=harm.dataset_name) for harm in _SUB_HARMS], ), # No single scenario objective scorer -- each sub-harm scores itself. The base contract # still requires one for scenario identity; the imminent-crisis scorer stands in. @@ -478,93 +474,53 @@ def _selected_sub_harms(self) -> list[_SubHarm]: ) return [_SUB_HARMS_BY_NAME[name]] - async def _resolve_seed_groups_by_dataset_async( - self, *, apply_sampling: bool = True - ) -> dict[str, list[AttackSeedGroup]]: - """ - Hard-bind the dataset names to the selected sub-harms before resolving seeds. - - Forces the dataset names to the selected sub-harms' (so ``--dataset-names`` cannot repoint - the scenario at unrelated data) and applies any ``max_dataset_size`` budget *per sub-harm* - rather than as one global budget. A single ``DatasetAttackConfiguration`` spends one budget - across the union of sub-harm datasets, so a small cap (e.g. ``1``) can starve a sub-harm of - every seed; a per-sub-harm compound caps each independently. Any run-time ``filters`` on the - active ``dataset_config`` are preserved. - - Args: - apply_sampling (bool): When True (default), apply ``max_dataset_size`` sampling. On - resume the base passes False so the full deterministic dataset is resolved. - - Returns: - dict[str, list[AttackSeedGroup]]: Seed groups keyed by originating dataset name. - """ - dataset_names = [harm.dataset_name for harm in self._selected_sub_harms()] - per_subharm_cap = self._dataset_config.max_dataset_size - filters = self._dataset_config.filters - if per_subharm_cap is None: - self._dataset_config = DatasetAttackConfiguration( - dataset_names=dataset_names, max_dataset_size=None, filters=filters - ) - else: - rebuilt = CompoundDatasetAttackConfiguration.per_dataset( - dataset_names=dataset_names, max_dataset_size=per_subharm_cap, filters=filters - ) - # Parent cap = per-sub-harm cap x sub-harm count: each child already caps at the - # per-sub-harm budget so the parent never trims the union, yet it stays non-None so the - # base still pins the sampled objective subset into the scenario metadata for resume. - rebuilt.max_dataset_size = per_subharm_cap * len(dataset_names) - self._dataset_config = rebuilt - return await super()._resolve_seed_groups_by_dataset_async(apply_sampling=apply_sampling) + def _validate_runtime_configuration(self) -> None: + config = self._dataset_config + allowed = {harm.dataset_name for harm in _SUB_HARMS} + if set(config.dataset_names) - allowed: + raise DatasetConstraintError("Psychosocial sources must match its sub-harm datasets.") + by_name = {source.name: source for source in config.sources} + self._dataset_config = config.with_overrides( + sources=[ + by_name.get(harm.dataset_name, DatasetSource(name=harm.dataset_name)) + for harm in self._selected_sub_harms() + ] + ) + super()._validate_runtime_configuration() def _get_run_size_budget(self) -> ScenarioDatasetSizeEstimate: """ - Use the outer limit applied by seed resolution, not compound child budgets. + Use the source and total limits applied by seed resolution. Returns: ScenarioDatasetSizeEstimate: Combined sub-harm cap or all available finite data. """ - cap = self._dataset_config.max_dataset_size - return ( - AllAvailableDatasetSize() - if cap is None - else BoundedDatasetSize(value=cap * len(self._selected_sub_harms())) - ) + return self._dataset_config.get_size_budget() async def _estimate_run_size_async(self, *, budget: BoundedDatasetSize) -> ScenarioRunSizeEstimate: """ Estimate the independent sub-harm technique sweeps and per-harm baselines. Returns: - ScenarioRunSizeEstimate: Configured per-sub-harm budget. + ScenarioRunSizeEstimate: Combined selected-population budget. """ - per_harm_count = budget.value // len(self._selected_sub_harms()) - datasets = [ - ScenarioDatasetSummary( - name=harm.dataset_name, - configured_caps=[ScenarioDatasetSizeCap(label="per-sub-harm cap", count=per_harm_count)], + count, datasets = await self._get_dataset_size_for_estimate_async(budget=budget) + technique_count = len(self._scenario_techniques) + components = [ + ScenarioRunSizeComponent( + label="Sub-harm technique sweeps", + count=count * technique_count, ) - for harm in self._selected_sub_harms() ] - technique_count = len(self._scenario_techniques) - components: list[ScenarioRunSizeComponent] = [] - for dataset in datasets: - dataset_name = dataset.name - count = per_harm_count + if self._include_baseline: components.append( ScenarioRunSizeComponent( - label=f"{dataset_name} technique sweep", - count=count * technique_count, + label="Sub-harm baselines", + count=count, + is_baseline=True, + note="Each selected objective has one baseline with its sub-harm scorer.", ) ) - if self._include_baseline: - components.append( - ScenarioRunSizeComponent( - label=f"{dataset_name} baseline", - count=count, - is_baseline=True, - note="Psychosocial uses a distinct baseline and scorer for each sub-harm.", - ) - ) return ScenarioRunSizeEstimate( total_attack_count=sum(component.count for component in components), components=components, diff --git a/pyrit/scenario/scenarios/airt/rapid_response.py b/pyrit/scenario/scenarios/airt/rapid_response.py index 4fd292bbe8..54a391de84 100644 --- a/pyrit/scenario/scenarios/airt/rapid_response.py +++ b/pyrit/scenario/scenarios/airt/rapid_response.py @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING from pyrit.common import apply_defaults -from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_matrix_atomic_attacks from pyrit.scenario.core.scenario import Scenario @@ -94,17 +94,20 @@ def __init__( version=self.VERSION, objective_scorer=self._objective_scorer, technique_class=technique_class, - default_dataset_config=CompoundDatasetAttackConfiguration.per_dataset( - dataset_names=[ - "airt_hate", - "airt_fairness", - "airt_violence", - "airt_sexual", - "airt_harassment", - "airt_misinformation", - "airt_leakage", + default_dataset_config=DatasetAttackConfiguration( + sources=[ + DatasetSource(name=name) + for name in [ + "airt_hate", + "airt_fairness", + "airt_violence", + "airt_sexual", + "airt_harassment", + "airt_misinformation", + "airt_leakage", + ] ], - max_dataset_size=4, + max_per_dataset=4, ), scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/airt/scam.py b/pyrit/scenario/scenarios/airt/scam.py index 686e277b24..dfd75adf95 100644 --- a/pyrit/scenario/scenarios/airt/scam.py +++ b/pyrit/scenario/scenarios/airt/scam.py @@ -15,7 +15,7 @@ from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack from pyrit.scenario.core.scenario import Scenario from pyrit.scenario.core.scenario_context import ScenarioContext @@ -153,7 +153,9 @@ def __init__( version=self.VERSION, uses_default_adversarial_target=adversarial_chat is None, technique_class=ScamTechnique, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["airt_scams"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["airt_scams"]], max_per_dataset="all", max_total=4 + ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/benchmark/adversarial.py b/pyrit/scenario/scenarios/benchmark/adversarial.py index 0daaf8cfe9..4f6bf7ee11 100644 --- a/pyrit/scenario/scenarios/benchmark/adversarial.py +++ b/pyrit/scenario/scenarios/benchmark/adversarial.py @@ -35,7 +35,7 @@ from pyrit.models.identifiers import compute_inner_attack_eval_hash from pyrit.models.parameter import Parameter from pyrit.registry import AttackTechniqueRegistry, TargetRegistry -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import ( MatrixAtomicAttackBuilder, resolve_technique_factories, @@ -262,8 +262,9 @@ def __init__( objective_scorer=self._objective_scorer, technique_class=technique_class, default_dataset_config=DatasetAttackConfiguration( - dataset_names=["harmbench"], - max_dataset_size=8, + sources=[DatasetSource(name=name) for name in ["harmbench"]], + max_per_dataset="all", + max_total=8, ), scenario_result_id=scenario_result_id, ) @@ -284,9 +285,19 @@ async def _resolve_seed_groups_by_dataset_async( return await super()._resolve_seed_groups_by_dataset_async(apply_sampling=apply_sampling) groups_by_dataset = await super()._resolve_seed_groups_by_dataset_async(apply_sampling=False) - max_dataset_size = self._dataset_config.max_dataset_size + if self._dataset_config.sources: + for name, groups in groups_by_dataset.items(): + limit = self._dataset_config.source_limit(name) + if limit != "all" and len(groups) > limit: + groups_by_dataset[name] = [ + group + for _, group in self._select_stable_sample( + pairs=[(name, group) for group in groups], max_dataset_size=limit + ) + ] + max_dataset_size = self._dataset_config.max_total pairs = [(name, group) for name, groups in groups_by_dataset.items() for group in groups] - if max_dataset_size is None or len(pairs) <= max_dataset_size: + if max_dataset_size == "all" or len(pairs) <= max_dataset_size: return groups_by_dataset selected = self._select_stable_sample(pairs=pairs, max_dataset_size=max_dataset_size) @@ -377,7 +388,9 @@ def _get_run_size_budget(self) -> ScenarioDatasetSizeEstimate: Returns: ScenarioDatasetSizeEstimate: Global cap or all available benchmark data. """ - return scenario_dataset_size_from_limit(self._dataset_config.max_dataset_size) + if self._dataset_config.sources: + return self._dataset_config.get_size_budget() + return scenario_dataset_size_from_limit(self._dataset_config.max_total) def _get_estimate_dataset_configuration(self) -> DatasetAttackConfiguration: """ @@ -386,9 +399,12 @@ def _get_estimate_dataset_configuration(self) -> DatasetAttackConfiguration: Returns: DatasetAttackConfiguration: Cap metadata without ignored child limits. """ + if self._dataset_config.sources: + return self._dataset_config return DatasetAttackConfiguration( - dataset_names=self._dataset_config.dataset_names or None, - max_dataset_size=self._dataset_config.max_dataset_size, + sources=[DatasetSource(name=name) for name in self._dataset_config.dataset_names], + max_per_dataset="all", + max_total=self._dataset_config.max_total, ) async def _estimate_run_size_async(self, *, budget: BoundedDatasetSize) -> ScenarioRunSizeEstimate: diff --git a/pyrit/scenario/scenarios/foundry/red_team_agent.py b/pyrit/scenario/scenarios/foundry/red_team_agent.py index 1b46e26183..f7230ce2ec 100644 --- a/pyrit/scenario/scenarios/foundry/red_team_agent.py +++ b/pyrit/scenario/scenarios/foundry/red_team_agent.py @@ -60,7 +60,7 @@ from pyrit.prompt_target import PromptTarget from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack from pyrit.scenario.core.scenario import Scenario from pyrit.scenario.core.scenario_context import ScenarioContext @@ -344,7 +344,9 @@ def __init__( version=self.VERSION, uses_default_adversarial_target=adversarial_chat is None, technique_class=FoundryTechnique, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=4), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["harmbench"]], max_per_dataset="all", max_total=4 + ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/garak/_prompt_injection.py b/pyrit/scenario/scenarios/garak/_prompt_injection.py index 96862df2f3..0b7c1d6579 100644 --- a/pyrit/scenario/scenarios/garak/_prompt_injection.py +++ b/pyrit/scenario/scenarios/garak/_prompt_injection.py @@ -8,6 +8,7 @@ from typing import TypeVar from pyrit.models import AttackSeedGroup +from pyrit.models.dataset_limit import ResolvedDatasetLimit from pyrit.scenario.core.dataset_configuration import DatasetConstraintError _Key = TypeVar("_Key", bound=Hashable) @@ -16,7 +17,7 @@ def sample_with_coverage( *, groups_by_dataset: dict[str, list[AttackSeedGroup]], - cap: int | None, + cap: ResolvedDatasetLimit, required_keys: Sequence[_Key], key: Callable[[AttackSeedGroup], _Key], ) -> dict[str, list[AttackSeedGroup]]: @@ -29,7 +30,7 @@ def sample_with_coverage( Raises: DatasetConstraintError: If a coverage key is missing or the budget is too small. """ - if cap is not None and cap < max(1, len(required_keys)): + if cap != "all" and cap < max(1, len(required_keys)): raise DatasetConstraintError( f"max_dataset_size ({cap}) must be at least the number of coverage groups ({len(required_keys)})." ) @@ -40,7 +41,7 @@ def sample_with_coverage( missing = [value for value, indices in indices_by_key.items() if not indices] if missing: raise DatasetConstraintError(f"No contexts for coverage groups: {missing}.") - if cap is None or len(pairs) <= cap: + if cap == "all" or len(pairs) <= cap: return groups_by_dataset selected = {random.choice(indices) for indices in indices_by_key.values()} remaining = [index for index in range(len(pairs)) if index not in selected] diff --git a/pyrit/scenario/scenarios/garak/api_key.py b/pyrit/scenario/scenarios/garak/api_key.py index 9409513daf..324858d3cf 100644 --- a/pyrit/scenario/scenarios/garak/api_key.py +++ b/pyrit/scenario/scenarios/garak/api_key.py @@ -31,6 +31,7 @@ from pyrit.scenario.core.dataset_configuration import ( DatasetAttackConfiguration, DatasetConstraintError, + DatasetSource, ResolvedDataset, ) from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario @@ -83,14 +84,20 @@ class ApiKeyDatasetConfiguration(DatasetAttackConfiguration): DEFAULT_MAX_DATASET_SIZE: ClassVar[int] = 20 @forward_init_parameters - def __init__(self, **kwargs: Any) -> None: + def __init__(self, *, sampling_scope: Literal["total_only"] = "total_only", **kwargs: Any) -> None: """ Initialize the configuration. Args: + sampling_scope (Literal["total_only"]): Sample assembled groups, not ingredient rows. **kwargs (Any): Arguments for ``DatasetAttackConfiguration``. + + Raises: + DatasetConstraintError: If sampling_scope is not total_only. """ - super().__init__(**kwargs) + if sampling_scope != "total_only": + raise DatasetConstraintError("ApiKey requires total_only sampling.") + super().__init__(sampling_scope=sampling_scope, **kwargs) self._techniques: list[ApiKeyTechnique] = [ApiKeyTechnique.GetKey, ApiKeyTechnique.CompleteKey] self.excluded_values: tuple[str, ...] = () @@ -105,10 +112,10 @@ def size_caps_by_dataset(self) -> dict[str, list[tuple[str, int, Literal["datase Returns: dict: Technique names mapped to their shared configuration cap. """ - if self.max_dataset_size is None: + if self.max_total == "all": return {} return { - str(technique.value): [("combined configuration cap", self.max_dataset_size, "configuration")] + str(technique.value): [("combined configuration cap", self.max_total, "configuration")] for technique in self._techniques } @@ -229,8 +236,9 @@ def __init__( version=self.VERSION, technique_class=ApiKeyTechnique, default_dataset_config=ApiKeyDatasetConfiguration( - dataset_names=self.required_datasets(), - max_dataset_size=ApiKeyDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, + sources=[DatasetSource(name=name) for name in self.required_datasets()], + max_per_dataset="all", + max_total=ApiKeyDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, ), objective_scorer=objective_scorer or CredentialLeakScorer(patterns=CredentialLeakScorer.GARAK_PATTERNS), scenario_result_id=scenario_result_id, diff --git a/pyrit/scenario/scenarios/garak/audio_achilles_heel.py b/pyrit/scenario/scenarios/garak/audio_achilles_heel.py index b332b7a468..9f5d9a48e2 100644 --- a/pyrit/scenario/scenarios/garak/audio_achilles_heel.py +++ b/pyrit/scenario/scenarios/garak/audio_achilles_heel.py @@ -32,7 +32,7 @@ from pyrit.models import AttackSeedGroup, Seed, SeedObjective, SeedPrompt from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique @@ -206,9 +206,10 @@ def __init__( version=self.VERSION, technique_class=AudioAchillesHeelTechnique, default_dataset_config=AudioAchillesHeelDatasetConfiguration( - dataset_names=["garak_audio_achilles_heel"], + sources=[DatasetSource(name=name) for name in ["garak_audio_achilles_heel"]], text_prompt=text_prompt, - max_dataset_size=DEFAULT_MAX_DATASET_SIZE, + max_per_dataset="all", + max_total=DEFAULT_MAX_DATASET_SIZE, ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, diff --git a/pyrit/scenario/scenarios/garak/divergence.py b/pyrit/scenario/scenarios/garak/divergence.py index 7305fbad1c..201df2866c 100644 --- a/pyrit/scenario/scenarios/garak/divergence.py +++ b/pyrit/scenario/scenarios/garak/divergence.py @@ -17,7 +17,7 @@ from pyrit.prompt_normalizer import ConverterConfiguration from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetConstraintError +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetConstraintError, DatasetSource from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.score import DivergenceScorer, TrueFalseScorer @@ -122,8 +122,9 @@ def __init__( version=self.VERSION, technique_class=DivergenceTechnique, default_dataset_config=DivergenceDatasetConfiguration( - dataset_names=self.required_datasets(), - max_dataset_size=DivergenceDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, + sources=[DatasetSource(name=name) for name in self.required_datasets()], + max_per_dataset="all", + max_total=DivergenceDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, ), objective_scorer=objective_scorer or DivergenceScorer(), scenario_result_id=scenario_result_id, diff --git a/pyrit/scenario/scenarios/garak/doctor.py b/pyrit/scenario/scenarios/garak/doctor.py index 1aed3279a9..8bb073d7b5 100644 --- a/pyrit/scenario/scenarios/garak/doctor.py +++ b/pyrit/scenario/scenarios/garak/doctor.py @@ -13,7 +13,7 @@ from pyrit.prompt_normalizer import ConverterConfiguration from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import MatrixAtomicAttackBuilder from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario @@ -137,7 +137,9 @@ def __init__( super().__init__( version=self.VERSION, technique_class=technique_class, - default_dataset_config=DatasetAttackConfiguration(dataset_names=["garak_doctor"]), + default_dataset_config=DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in ["garak_doctor"]] + ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/garak/encoding.py b/pyrit/scenario/scenarios/garak/encoding.py index e1a95d9108..5c57600510 100644 --- a/pyrit/scenario/scenarios/garak/encoding.py +++ b/pyrit/scenario/scenarios/garak/encoding.py @@ -36,7 +36,11 @@ from pyrit.prompt_normalizer.converter_configuration import ConverterConfiguration from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration, DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import ( + CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, + DatasetSource, +) from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack from pyrit.scenario.core.scenario import Scenario from pyrit.scenario.core.scenario_context import ScenarioContext @@ -196,8 +200,16 @@ def __init__( technique_class=EncodingTechnique, default_dataset_config=CompoundDatasetAttackConfiguration( configurations=[ - EncodingDatasetConfiguration(dataset_names=["garak_slur_terms_en"], max_dataset_size=10), - EncodingDatasetConfiguration(dataset_names=["garak_web_html_js"], max_dataset_size=10), + EncodingDatasetConfiguration( + sources=[DatasetSource(name=name) for name in ["garak_slur_terms_en"]], + max_per_dataset="all", + max_total=10, + ), + EncodingDatasetConfiguration( + sources=[DatasetSource(name=name) for name in ["garak_web_html_js"]], + max_per_dataset="all", + max_total=10, + ), ] ), objective_scorer=objective_scorer, diff --git a/pyrit/scenario/scenarios/garak/exploitation.py b/pyrit/scenario/scenarios/garak/exploitation.py index 2b7235c067..66a500699e 100644 --- a/pyrit/scenario/scenarios/garak/exploitation.py +++ b/pyrit/scenario/scenarios/garak/exploitation.py @@ -31,7 +31,7 @@ from pyrit.prompt_normalizer import ConverterConfiguration from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.score import ( @@ -253,7 +253,9 @@ def __init__( super().__init__( version=self.VERSION, technique_class=ExploitationTechnique, - default_dataset_config=_ExploitationDatasetConfiguration(dataset_names=list(_CORPUS_DATASETS)), + default_dataset_config=_ExploitationDatasetConfiguration( + sampling_scope="total_only", sources=[DatasetSource(name=name) for name in list(_CORPUS_DATASETS)] + ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, ) diff --git a/pyrit/scenario/scenarios/garak/figstep.py b/pyrit/scenario/scenarios/garak/figstep.py index 3e580fd165..3dea0f7a3b 100644 --- a/pyrit/scenario/scenarios/garak/figstep.py +++ b/pyrit/scenario/scenarios/garak/figstep.py @@ -19,6 +19,7 @@ INLINE_DATASET_NAME, DatasetAttackConfiguration, DatasetConfiguration, + DatasetSource, DatasetSourceKind, ) from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack @@ -103,8 +104,9 @@ def __init__( version=self.VERSION, technique_class=FigStepTechnique, default_dataset_config=DatasetAttackConfiguration( - dataset_names=["figstep"], - max_dataset_size=DEFAULT_MAX_DATASET_SIZE, + sources=[DatasetSource(name=name) for name in ["figstep"]], + max_per_dataset="all", + max_total=DEFAULT_MAX_DATASET_SIZE, ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, @@ -169,7 +171,7 @@ async def _resolve_seed_groups_by_dataset_async( Returns: dict[str, list[AttackSeedGroup]]: Valid FigStep groups keyed by dataset. """ - validate_before_sampling = apply_sampling and self._dataset_config.max_dataset_size is not None + validate_before_sampling = apply_sampling and self._dataset_config.has_sampling_limits validation_sampling = apply_sampling and not validate_before_sampling groups_by_dataset = await super()._resolve_seed_groups_by_dataset_async(apply_sampling=validation_sampling) self._validate_seed_groups(seed_groups=[group for groups in groups_by_dataset.values() for group in groups]) diff --git a/pyrit/scenario/scenarios/garak/latent_injection.py b/pyrit/scenario/scenarios/garak/latent_injection.py index 175e5e0d4f..3c6ae9dd5e 100644 --- a/pyrit/scenario/scenarios/garak/latent_injection.py +++ b/pyrit/scenario/scenarios/garak/latent_injection.py @@ -12,7 +12,7 @@ import itertools import json import re -from typing import TYPE_CHECKING, Any, ClassVar, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast from pyrit.common import apply_defaults, forward_init_parameters from pyrit.converter import SearchReplaceConverter @@ -24,6 +24,7 @@ from pyrit.scenario.core.dataset_configuration import ( DatasetAttackConfiguration, DatasetConstraintError, + DatasetSource, ResolvedDataset, ) from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario @@ -104,7 +105,7 @@ def __init__( self, *, families: Sequence[str] | None = None, - max_dataset_size: int | None = DEFAULT_MAX_DATASET_SIZE, + sampling_scope: Literal["total_only"] = "total_only", **kwargs: Any, ) -> None: """ @@ -112,13 +113,21 @@ def __init__( Args: families (Sequence[str] | None): Selected families, excluding latent jailbreak by default. - max_dataset_size (int | None): Maximum selected groups. Defaults to 92; None selects all groups. + sampling_scope (Literal["total_only"]): Sample assembled groups, not ingredient rows. **kwargs (Any): Standard dataset settings. An explicit uncapped configuration uses all groups. + + Raises: + DatasetConstraintError: If sampling_scope is not total_only. """ - super().__init__(max_dataset_size=max_dataset_size, **kwargs) + if sampling_scope != "total_only": + raise DatasetConstraintError("LatentInjection requires total_only sampling.") + super().__init__(sampling_scope=sampling_scope, **kwargs) self._set_families(families=self.DEFAULT_FAMILIES if families is None else families) self.coverage_keys: list[tuple[str, str]] = [] + def _default_max_total(self) -> int: + return self.DEFAULT_MAX_DATASET_SIZE + @property def families(self) -> list[str]: """The selected families in declaration order.""" @@ -176,8 +185,8 @@ async def _build_groups_by_dataset_async(self) -> tuple[dict[str, list[AttackSee group for family in self._families for group in self._build_family_groups(family=family, roles=by_role) ] self.coverage_keys = sorted({self._coverage_key(group) for group in groups}) - cap = self.max_dataset_size - if cap is not None and cap < max(1, len(self.coverage_keys)): + cap = self.max_total + if cap != "all" and cap < max(1, len(self.coverage_keys)): raise DatasetConstraintError( f"max_dataset_size ({cap}) must be at least the number of family/trigger pairs " f"({len(self.coverage_keys)})." @@ -291,7 +300,7 @@ def _sample_groups_by_dataset( ) -> dict[str, list[AttackSeedGroup]]: return sample_with_coverage( groups_by_dataset=groups_by_dataset, - cap=self.max_dataset_size, + cap=self.max_total, required_keys=self.coverage_keys, key=self._coverage_key, ) @@ -346,8 +355,9 @@ def __init__( version=self.VERSION, technique_class=LatentInjectionTechnique, default_dataset_config=LatentInjectionDatasetConfiguration( - dataset_names=self.required_datasets(), - max_dataset_size=LatentInjectionDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, + sources=[DatasetSource(name=name) for name in self.required_datasets()], + max_per_dataset="all", + max_total=LatentInjectionDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, ), objective_scorer=objective_scorer or SubStringScorer(substring="Haha pwned!"), scenario_result_id=scenario_result_id, diff --git a/pyrit/scenario/scenarios/garak/package_hallucination.py b/pyrit/scenario/scenarios/garak/package_hallucination.py index b4d9b0c97f..80be685a74 100644 --- a/pyrit/scenario/scenarios/garak/package_hallucination.py +++ b/pyrit/scenario/scenarios/garak/package_hallucination.py @@ -25,7 +25,7 @@ ) from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.score.true_false.regex.package_hallucination_scorer import ( @@ -88,20 +88,6 @@ class _LanguageSpec: } -class _PackageHallucinationDatasetConfiguration(DatasetConfiguration): - """Dataset configuration that exposes raw values for prompt and registry datasets.""" - - async def get_values_by_dataset_async(self) -> dict[str, list[str]]: - """ - Resolve configured datasets, fetching missing datasets from their providers. - - Returns: - dict[str, list[str]]: Seed values keyed by dataset name. - """ - seeds_by_dataset = await self._collect_named_seeds_async() - return {name: [seed.value for seed in seeds] for name, seeds in seeds_by_dataset.items()} - - class PackageHallucinationTechnique(ScenarioTechnique): """ Techniques for the PackageHallucination scenario. @@ -205,7 +191,10 @@ def __init__( # Preload only the Rust registry and prompt corpus. Other registries are fetched # on demand when their techniques are selected. default_dataset_config=DatasetAttackConfiguration( - dataset_names=[_LANGUAGE_SPECS["rust"].dataset_name, *_CORPUS_DATASETS] + sampling_scope="total_only", + sources=[ + DatasetSource(name=name) for name in [_LANGUAGE_SPECS["rust"].dataset_name, *_CORPUS_DATASETS] + ], ), objective_scorer=objective_scorer, scenario_result_id=scenario_result_id, @@ -214,6 +203,14 @@ def __init__( USES_DATASET_SIZE_LIMIT: ClassVar[bool] = False def _validate_runtime_configuration(self) -> None: + names = [ + *_CORPUS_DATASETS, + *(_LANGUAGE_SPECS[technique.value].dataset_name for technique in self._scenario_techniques), + ] + by_name = {source.name: source for source in self._dataset_config.sources} + self._dataset_config = self._dataset_config.with_overrides( + sources=[by_name.get(name, DatasetSource(name=name)) for name in dict.fromkeys(names)] + ) super()._validate_runtime_configuration() if self._max_prompts_per_language < 1: raise ValueError("max_prompts_per_language must be greater than zero") @@ -342,13 +339,8 @@ async def _resolve_seed_groups_by_dataset_async( if not isinstance(technique, PackageHallucinationTechnique): raise TypeError(f"Unexpected package hallucination technique: {type(technique).__name__}") specs_by_technique[technique.value] = _LANGUAGE_SPECS[technique.value] - dataset_names = [ - *_CORPUS_DATASETS, - *(spec.dataset_name for spec in specs_by_technique.values()), - ] - dataset_values = await _PackageHallucinationDatasetConfiguration( - dataset_names=list(dict.fromkeys(dataset_names)) - ).get_values_by_dataset_async() + seeds_by_dataset = await self._dataset_config._collect_named_seeds_async() + dataset_values = {name: [seed.value for seed in seeds] for name, seeds in seeds_by_dataset.items()} rng = random.Random(self._random_seed) stubs, tasks = self._load_corpus(dataset_values=dataset_values) diff --git a/pyrit/scenario/scenarios/garak/prompt_inject.py b/pyrit/scenario/scenarios/garak/prompt_inject.py index 754c6b7d83..c43e08e263 100644 --- a/pyrit/scenario/scenarios/garak/prompt_inject.py +++ b/pyrit/scenario/scenarios/garak/prompt_inject.py @@ -10,7 +10,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, ClassVar, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast from pyrit.common import apply_defaults, forward_init_parameters from pyrit.converter import SearchReplaceConverter @@ -20,7 +20,11 @@ from pyrit.prompt_normalizer import ConverterConfiguration from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetConstraintError +from pyrit.scenario.core.dataset_configuration import ( + DatasetAttackConfiguration, + DatasetConstraintError, + DatasetSource, +) from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.scenario.scenarios.garak._prompt_injection import sample_with_coverage @@ -58,7 +62,7 @@ def __init__( self, *, goal_texts: Sequence[str] | None = None, - max_dataset_size: int | None = DEFAULT_MAX_DATASET_SIZE, + sampling_scope: Literal["total_only"] = "total_only", **kwargs: Any, ) -> None: """ @@ -66,16 +70,22 @@ def __init__( Args: goal_texts (Sequence[str] | None): Text that the target is asked to return. - max_dataset_size (int | None): Maximum selected groups. Defaults to 12; None selects all groups. + sampling_scope (Literal["total_only"]): Sample assembled groups, not ingredient rows. **kwargs (Any): Arguments for ``DatasetAttackConfiguration``. Raises: ValueError: If goal texts are empty or duplicated. + DatasetConstraintError: If sampling_scope is not total_only. """ - super().__init__(max_dataset_size=max_dataset_size, **kwargs) + if sampling_scope != "total_only": + raise DatasetConstraintError("PromptInject requires total_only sampling.") + super().__init__(sampling_scope=sampling_scope, **kwargs) goal_texts = _DEFAULT_GOAL_TEXTS if goal_texts is None else goal_texts self._set_goal_texts(goal_texts=goal_texts) + def _default_max_total(self) -> int: + return self.DEFAULT_MAX_DATASET_SIZE + def _set_goal_texts(self, *, goal_texts: Sequence[str]) -> None: """ Set the goal texts used to build attack groups. @@ -113,7 +123,7 @@ def _sample_groups_by_dataset( """ return sample_with_coverage( groups_by_dataset=groups_by_dataset, - cap=self.max_dataset_size, + cap=self.max_total, required_keys=self._goal_texts, key=lambda group: (group.objective.metadata or {})["goal_text"], ) @@ -135,8 +145,8 @@ def validate_configuration(self) -> None: raise DatasetConstraintError( f"PromptInject requires exactly these datasets: {sorted(required_dataset_names)}." ) - cap = self.max_dataset_size - if cap is not None and cap < len(self._goal_texts): + cap = self.max_total + if cap != "all" and cap < len(self._goal_texts): raise DatasetConstraintError( f"PromptInject max_dataset_size ({cap}) must be at least the number of goal_texts " f"({len(self._goal_texts)})." @@ -259,8 +269,9 @@ def __init__( version=self.VERSION, technique_class=PromptInjectTechnique, default_dataset_config=PromptInjectDatasetConfiguration( - dataset_names=self.required_datasets(), - max_dataset_size=PromptInjectDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, + sources=[DatasetSource(name=name) for name in self.required_datasets()], + max_per_dataset="all", + max_total=PromptInjectDatasetConfiguration.DEFAULT_MAX_DATASET_SIZE, goal_texts=self.DEFAULT_GOAL_TEXTS, ), objective_scorer=objective_scorer, diff --git a/pyrit/scenario/scenarios/garak/system_prompt_extraction.py b/pyrit/scenario/scenarios/garak/system_prompt_extraction.py index cb8ff5e4c4..519dcc6b3d 100644 --- a/pyrit/scenario/scenarios/garak/system_prompt_extraction.py +++ b/pyrit/scenario/scenarios/garak/system_prompt_extraction.py @@ -21,9 +21,10 @@ SeedPrompt, scenario_dataset_size_from_limit, ) +from pyrit.models.dataset_limit import DatasetLimit, normalize_dataset_limit from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.score import FloatScaleThresholdScorer, SystemPromptExtractionScorer @@ -113,7 +114,7 @@ def __init__( *, objective_scorer: TrueFalseScorer | None = None, system_prompt_subsample: int = 50, - prompt_cap: int | None = _DEFAULT_PROMPT_CAP, + prompt_cap: DatasetLimit = "default", random_seed: int | None = None, scenario_result_id: str | None = None, ) -> None: @@ -126,13 +127,16 @@ def __init__( ``SystemPromptExtractionScorer`` (n=4) at threshold 0.5 (garak's ``eval_threshold``). system_prompt_subsample (int): Maximum number of system prompts to draw per dataset. Defaults to 50 (garak's ``system_prompt_subsample``). - prompt_cap (int | None): Upper bound on the total number of (system prompt x template) + prompt_cap (DatasetLimit): Upper bound on the total number of (system prompt x template) sends per run. The full combination set is randomly sampled down to this size, - mirroring garak's ``soft_probe_prompt_cap``. Set to None to run every combination. - Defaults to 256. + mirroring garak's ``soft_probe_prompt_cap``. Use "all" for every combination. + Omitted, None, empty, and "default" use 256. random_seed (int | None): Seed for deterministic sampling of system prompts and the prompt cap. Defaults to a fixed value for reproducibility. scenario_result_id (str | None): Optional ID of an existing scenario result to resume. + + Raises: + ValueError: If prompt_cap is not a supported dataset limit. """ if not objective_scorer: objective_scorer = FloatScaleThresholdScorer( @@ -141,17 +145,22 @@ def __init__( ) self._scorer_config = AttackScoringConfig(objective_scorer=objective_scorer) self._system_prompt_subsample = system_prompt_subsample - self._prompt_cap = prompt_cap + limit = normalize_dataset_limit(prompt_cap) + self._prompt_cap = _DEFAULT_PROMPT_CAP if limit == "default" else limit self._random_seed = random_seed if random_seed is not None else 42 super().__init__( version=self.VERSION, technique_class=SystemPromptExtractionTechnique, default_dataset_config=DatasetAttackConfiguration( - dataset_names=[ - DATASET_DRH_SYSTEM_PROMPTS, - DATASET_TM_SYSTEM_PROMPTS, - DATASET_EXTRACTION_TEMPLATES, + max_per_dataset="all", + sources=[ + DatasetSource(name=name) + for name in [ + DATASET_DRH_SYSTEM_PROMPTS, + DATASET_TM_SYSTEM_PROMPTS, + DATASET_EXTRACTION_TEMPLATES, + ] ], ), objective_scorer=objective_scorer, @@ -162,8 +171,6 @@ def __init__( def _validate_runtime_configuration(self) -> None: super()._validate_runtime_configuration() - if self._prompt_cap is not None and self._prompt_cap < 1: - raise ValueError("prompt_cap must be greater than zero or None") if self._system_prompt_subsample < 1: raise ValueError("system_prompt_subsample must be greater than zero") @@ -264,7 +271,7 @@ async def _resolve_seed_groups_by_dataset_async( f"{DATASET_EXTRACTION_TEMPLATES}) are loaded into CentralMemory before running." ) - if self._prompt_cap is not None and len(combinations) > self._prompt_cap: + if self._prompt_cap != "all" and len(combinations) > self._prompt_cap: combinations = random.Random(self._random_seed).sample(combinations, self._prompt_cap) seed_groups_by_category: dict[str, list[AttackSeedGroup]] = {} diff --git a/pyrit/scenario/scenarios/garak/web_injection.py b/pyrit/scenario/scenarios/garak/web_injection.py index e9e4ad450f..76551926bf 100644 --- a/pyrit/scenario/scenarios/garak/web_injection.py +++ b/pyrit/scenario/scenarios/garak/web_injection.py @@ -25,7 +25,7 @@ ) from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique -from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.core.matrix_atomic_attack_builder import build_baseline_atomic_attack from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario from pyrit.scenario.core.scenario_technique import ScenarioTechnique @@ -264,11 +264,15 @@ def __init__( version=self.VERSION, technique_class=WebInjectionTechnique, default_dataset_config=DatasetAttackConfiguration( - dataset_names=[ - self.DATASET_EXAMPLE_DOMAINS, - self.DATASET_MARKDOWN_JS, - self.DATASET_WEB_HTML_JS, - self.DATASET_NORMAL_INSTRUCTIONS, + max_per_dataset="all", + sources=[ + DatasetSource(name=name) + for name in [ + self.DATASET_EXAMPLE_DOMAINS, + self.DATASET_MARKDOWN_JS, + self.DATASET_WEB_HTML_JS, + self.DATASET_NORMAL_INSTRUCTIONS, + ] ], ), objective_scorer=objective_scorer, diff --git a/tests/unit/backend/test_scenario_configuration_resolver.py b/tests/unit/backend/test_scenario_configuration_resolver.py index c5fe5a3b5b..22ff1e2c0d 100644 --- a/tests/unit/backend/test_scenario_configuration_resolver.py +++ b/tests/unit/backend/test_scenario_configuration_resolver.py @@ -3,25 +3,34 @@ """Adversarial target resolution validates without changing execution scopes.""" -from unittest.mock import MagicMock, patch +from typing import Literal +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pyrit.backend.services.scenario_configuration_resolver import ScenarioConfigurationResolver +from pyrit.memory import CentralMemory, MemoryInterface +from pyrit.models import SeedObjective +from pyrit.models.dataset_limit import DatasetLimit, ResolvedDatasetLimit from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.registry import TargetRegistry +from pyrit.scenario import DatasetAttackConfiguration, DatasetFetchPolicy, DatasetSource, Scenario from pyrit.scenario.core import ( get_default_adversarial_target, override_default_adversarial_target, scenario_target_defaults, ) +from pyrit.scenario.scenarios.adaptive.text_adaptive import TextAdaptive +from pyrit.scenario.scenarios.airt.rapid_response import RapidResponse from pyrit.scenario.scenarios.garak.api_key import ApiKey +from pyrit.scenario.scenarios.garak.prompt_inject import PromptInject, PromptInjectDatasetConfiguration +from pyrit.score import TrueFalseScorer from unit.mocks import MockPromptTarget @pytest.mark.usefixtures("patch_central_database") -@pytest.mark.parametrize(("limit", "expected"), [(None, 20), (7, 7)]) -def test_dataset_name_override_preserves_scenario_default(*, limit: int | None, expected: int) -> None: +@pytest.mark.parametrize(("limit", "expected"), [(None, 20), ("", 20), ("default", 20), ("all", "all"), (7, 7)]) +def test_dataset_name_override_applies_explicit_total(*, limit: DatasetLimit, expected: ResolvedDatasetLimit) -> None: resolved = ScenarioConfigurationResolver.resolve_configuration( scenario_name="garak.api_key", scenario_class=ApiKey, @@ -31,6 +40,101 @@ def test_dataset_name_override_preserves_scenario_default(*, limit: int | None, assert resolved["dataset_config"].max_dataset_size == expected +@pytest.mark.usefixtures("patch_central_database") +def test_dataset_overrides_preserve_subclass_state_and_default() -> None: + scenario = PromptInject() + original = PromptInjectDatasetConfiguration( + dataset_names=PromptInject.required_datasets(), + goal_texts=["first goal", "second goal"], + ) + scenario._default_dataset_config = original + resolved = ScenarioConfigurationResolver.resolve_configuration( + scenario_name="garak.prompt_inject", + scenario_class=MagicMock(return_value=scenario), + dataset_names=PromptInject.required_datasets(), + max_dataset_size=7, + dataset_filters={"data_types": ["text"]}, + )["dataset_config"] + assert type(resolved) is PromptInjectDatasetConfiguration + assert resolved._goal_texts == ["first goal", "second goal"] + assert resolved.max_total == 7 + assert resolved.filters == {"data_types": ["text"]} + assert original.max_total == 12 + assert original.filters == {} + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("scenario_class", [RapidResponse, TextAdaptive]) +@pytest.mark.parametrize("total", [10, 100, "all"]) +async def test_total_override_preserves_source_limits_async( + *, scenario_class: type[Scenario], total: int | Literal["all"] +) -> None: + technique_class = PromptInject()._technique_class + with ( + patch.object(Scenario, "_get_default_objective_scorer", return_value=MagicMock(spec=TrueFalseScorer)), + patch( + "pyrit.scenario.scenarios.airt.rapid_response._build_rapid_response_technique", + return_value=technique_class, + ), + patch.object(TextAdaptive, "get_technique_class", return_value=technique_class), + ): + scenario = scenario_class() + original = scenario._default_dataset_config + resolved = ScenarioConfigurationResolver.resolve_configuration( + scenario_name="test", + scenario_class=MagicMock(return_value=scenario), + max_dataset_size=total, + )["dataset_config"] + populations = { + name: [SeedObjective(value=f"{name}-{index}", dataset_name=name) for index in range(20)] + for name in original.dataset_names + } + memory = MagicMock(spec=MemoryInterface) + memory.get_seeds_async = AsyncMock(side_effect=lambda *, dataset_name: populations[dataset_name]) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + groups = await resolved.get_attack_groups_by_dataset_async() + assert sum(len(population) for population in groups.values()) == ( + 4 * len(populations) if total == "all" else min(total, 4 * len(populations)) + ) + assert all(len(population) <= 4 for population in groups.values()) + assert resolved.max_total == total + assert resolved.max_per_dataset == 4 + assert original.max_total == "all" + assert original.max_per_dataset == 4 + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("total", [None, 7]) +def test_name_override_retains_source_options_and_adds_defaults_for_new_names(total: int | None) -> None: + retained = DatasetSource(name="retained", max_size=2, fetch=DatasetFetchPolicy.NEVER) + original = DatasetAttackConfiguration( + sources=[retained, DatasetSource(name="removed")], max_per_dataset=4, max_total=20 + ) + scenario = PromptInject() + scenario._default_dataset_config = original + resolved = ScenarioConfigurationResolver.resolve_configuration( + scenario_name="test", + scenario_class=MagicMock(return_value=scenario), + dataset_names=["new", "retained"], + max_dataset_size=total, + )["dataset_config"] + assert resolved.sources == (DatasetSource(name="new"), retained) + assert resolved.sources[1] is retained + assert resolved.source_limit("new") == 4 + assert resolved.source_limit("retained") == 2 + assert resolved.max_total == (20 if total is None else total) + assert original.dataset_names == ["retained", "removed"] + assert original.max_total == 20 + + +@pytest.mark.usefixtures("patch_central_database") +def test_omitted_total_retains_default() -> None: + resolved = ScenarioConfigurationResolver.resolve_configuration( + scenario_name="garak.api_key", scenario_class=ApiKey, dataset_names=ApiKey.required_datasets() + ) + assert resolved["dataset_config"].max_total == 20 + + @pytest.mark.usefixtures("patch_central_database") @pytest.mark.parametrize("selection", [None, "selected"]) def test_resolve_adversarial_target_does_not_mutate_scope(selection: str | None) -> None: diff --git a/tests/unit/backend/test_scenario_resume.py b/tests/unit/backend/test_scenario_resume.py index 472dc09042..a85af70e57 100644 --- a/tests/unit/backend/test_scenario_resume.py +++ b/tests/unit/backend/test_scenario_resume.py @@ -448,7 +448,9 @@ async def test_resume_missing_scenario_registration_is_explicit_async( await service.resume_run_async(scenario_result_id=str(stored.id)) -@pytest.mark.parametrize("missing", [name for name in _LAUNCH_REQUEST_FIELDS if name != "adversarial_target_name"]) +@pytest.mark.parametrize( + "missing", [name for name in _LAUNCH_REQUEST_FIELDS if name not in {"adversarial_target_name", "max_dataset_size"}] +) async def test_resume_incomplete_saved_configuration_never_uses_defaults_async( *, resume_environment: tuple[ScenarioRunService, MockPromptTarget], missing: str ) -> None: @@ -466,6 +468,33 @@ async def test_resume_incomplete_saved_configuration_never_uses_defaults_async( prepare.assert_not_called() +@pytest.mark.parametrize( + "limit_args", + [ + {}, + {"max_dataset_size": None}, + {"max_dataset_size": ""}, + {"max_dataset_size": "default"}, + {"max_dataset_size": "all"}, + {"max_dataset_size": 1}, + ], +) +async def test_saved_total_limit_round_trips_async( + *, resume_environment: tuple[ScenarioRunService, MockPromptTarget], limit_args: dict[str, int | str | None] +) -> None: + service, _ = resume_environment + prepared = await service._prepare_run_async( + request=RunScenarioRequest(scenario_name=_SCENARIO_NAME, target_name=_TARGET_NAME, **limit_args) + ) + stored = ( + await CentralMemory.get_memory_instance().get_scenario_results_async( + scenario_result_ids=[prepared.scenario._scenario_result_id] + ) + )[0] + restored = service._restore_launch_request(stored=stored) + assert restored.max_dataset_size == (limit_args.get("max_dataset_size") or "default") + + async def test_resume_older_launch_record_without_adversarial_selection_async( resume_environment: tuple[ScenarioRunService, MockPromptTarget], ) -> None: diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index cb12bf8f6a..a5803f2cc4 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -152,7 +152,7 @@ def _make_request( techniques=techniques, scenario_result_id=scenario_result_id, dataset_names=dataset_names, - max_dataset_size=max_dataset_size, + **({"max_dataset_size": max_dataset_size} if max_dataset_size is not None else {}), dataset_filters=dataset_filters, include_baseline=include_baseline, scenario_params=scenario_params, @@ -766,19 +766,44 @@ async def test_start_run_forwards_include_baseline(self, mock_all_registries) -> assert init_call.kwargs["include_baseline"] is False async def test_start_run_max_dataset_size_uses_default_config(self, mock_all_registries) -> None: - """``max_dataset_size`` with no ``dataset_names`` reuses the scenario's default config.""" - default_config = MagicMock() - default_config.max_dataset_size = 100 # original + """A total override copies the default configuration without mutating it.""" + default_config = DatasetAttackConfiguration(dataset_names=["original"], max_total=100) scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._default_dataset_config = default_config service = ScenarioRunService() await service.start_run_async(request=_make_request(max_dataset_size=5)) - # max_dataset_size on the default config was overridden - assert default_config.max_dataset_size == 5 + assert default_config.max_total == 100 init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args - assert init_call.kwargs["dataset_config"] is default_config + assert init_call.kwargs["dataset_config"] is not default_config + assert init_call.kwargs["dataset_config"].max_total == 5 + + @pytest.mark.parametrize( + ("limit_args", "expected"), + [ + ({}, 20), + ({"max_dataset_size": None}, 20), + ({"max_dataset_size": ""}, 20), + ({"max_dataset_size": "default"}, 20), + ({"max_dataset_size": "all"}, "all"), + ({"max_dataset_size": 7}, 7), + ], + ) + async def test_start_run_resolves_total_limit_async( + self, *, mock_all_registries: dict[str, Any], limit_args: dict[str, Any], expected: int | str + ) -> None: + default_config = DatasetAttackConfiguration(dataset_names=["original"], max_total=20) + mock_all_registries["scenario_instance"]._default_dataset_config = default_config + service = ScenarioRunService() + await service.start_run_async( + request=RunScenarioRequest(scenario_name="test", target_name="my_target", **limit_args) + ) + init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args + config = init_call.kwargs.get("dataset_config", default_config) + assert config.max_total == expected + saved = init_call.kwargs["initial_metadata"][_svc_mod._LAUNCH_REQUEST_METADATA_KEY] + assert saved["max_dataset_size"] == (limit_args.get("max_dataset_size") or "default") async def test_start_run_dataset_names_preserves_subclass_config_type(self, mock_all_registries) -> None: """``dataset_names`` rebuilds the config using the scenario's own DatasetConfiguration subclass. @@ -830,10 +855,12 @@ class _MarkerDatasetConfiguration(DatasetConfiguration): built_config = init_call.kwargs["dataset_config"] assert type(built_config) is _MarkerDatasetConfiguration assert built_config.dataset_names == ["only_this"] - assert built_config.max_dataset_size is None + assert built_config.max_dataset_size == "all" - async def test_start_run_dataset_names_rejects_incompatible_subclass_constructor(self, mock_all_registries) -> None: - """Reject overrides that cannot preserve scenario-specific dataset configuration.""" + async def test_start_run_dataset_names_preserves_custom_constructor_state_async( + self, mock_all_registries: dict[str, Any] + ) -> None: + """Copy overrides without calling a subclass constructor again.""" class _RequiresExtraArgConfiguration(DatasetConfiguration): def __init__(self, *, required_extra: str, **kwargs: Any) -> None: @@ -847,13 +874,12 @@ def __init__(self, *, required_extra: str, **kwargs: Any) -> None: ) service = ScenarioRunService() - with pytest.raises( - ValueError, - match="does not support overriding dataset names.*_RequiresExtraArgConfiguration", - ): - await service.start_run_async(request=_make_request(dataset_names=["custom"])) - - mock_all_registries["scenario_registry"].create_and_initialize_async.assert_not_awaited() + await service.start_run_async(request=_make_request(dataset_names=["custom"])) + init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args + config = init_call.kwargs["dataset_config"] + assert isinstance(config, _RequiresExtraArgConfiguration) + assert config._required_extra == "seeded" + assert config.dataset_names == ["custom"] async def test_start_run_dataset_filters_new_config(self, mock_all_registries) -> None: """``dataset_filters`` with ``dataset_names`` builds a config carrying the filters.""" @@ -890,7 +916,8 @@ async def test_start_run_dataset_filters_updates_default_config(self, mock_all_r init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args built_config = init_call.kwargs["dataset_config"] - assert built_config is default_config + assert built_config is not default_config + assert default_config.filters == {} assert built_config.filters == {"harm_categories": ["cyber"]} async def test_start_run_dataset_names_introspection_failure_raises(self, mock_memory) -> None: diff --git a/tests/unit/backend/test_scenario_service.py b/tests/unit/backend/test_scenario_service.py index 4b3551baa3..28edf854ef 100644 --- a/tests/unit/backend/test_scenario_service.py +++ b/tests/unit/backend/test_scenario_service.py @@ -514,7 +514,7 @@ async def test_estimate_is_offloaded_and_cached(self) -> None: assert second.default_run_size.datasets == estimate.datasets service._registry.create_instance.assert_called_once_with("test.scenario") - async def test_default_catalog_estimate_uses_read_only_dataset_resolution(self) -> None: + async def test_default_catalog_estimate_does_not_prepare_datasets(self) -> None: """Bulk catalog estimates do not auto-fetch datasets into memory.""" metadata = _make_scenario_metadata() estimate = ScenarioRunSizeEstimate( @@ -526,7 +526,7 @@ async def test_default_catalog_estimate_uses_read_only_dataset_resolution(self) with ( patch.object(ScenarioService, "__init__", lambda self: None), - patch("pyrit.backend.services.scenario_service.read_only_dataset_resolution") as read_only_resolution, + patch("pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async") as prepare, ): service = ScenarioService() service._registry = MagicMock() @@ -535,7 +535,7 @@ async def test_default_catalog_estimate_uses_read_only_dataset_resolution(self) result = await service._get_default_run_size_estimate_async(metadata=metadata) assert result == estimate - read_only_resolution.assert_called_once_with() + prepare.assert_not_awaited() async def test_concurrent_estimate_reads_share_one_task(self) -> None: """Concurrent catalog readers share one atomic single-flight estimate.""" @@ -1436,6 +1436,54 @@ async def test_list_scenarios_last_page_has_more_false(self) -> None: class TestScenarioServiceGetScenario: """Tests for ScenarioService.get_scenario_async.""" + @pytest.mark.parametrize( + ("limit_args", "expected"), + [ + ({}, 20), + ({"max_dataset_size": None}, 20), + ({"max_dataset_size": ""}, 20), + ({"max_dataset_size": "default"}, 20), + ({"max_dataset_size": "all"}, "all"), + ({"max_dataset_size": 7}, 7), + ], + ) + async def test_estimate_resolves_total_limit_async( + self, *, limit_args: dict[str, int | str | None], expected: int | str + ) -> None: + original = DatasetAttackConfiguration(dataset_names=["harmbench"], max_total=20) + scenario_class = MagicMock() + scenario_class.return_value._default_dataset_config = original + with patch.object(ScenarioService, "__init__", lambda self: None): + service = ScenarioService() + service._registry = MagicMock() + service._registry.get_registered_class_metadata.return_value = _make_scenario_metadata() + service._registry.get_class.return_value = scenario_class + service._registry.create_and_estimate_async = AsyncMock(return_value=ScenarioRunSizeEstimate()) + await service.estimate_scenario_run_size_async( + scenario_name="test.scenario", + request=ScenarioRunSizeEstimateRequest(dataset_names=["harmbench"], **limit_args), + ) + config = service._registry.create_and_estimate_async.await_args.kwargs["dataset_config"] + assert config.max_total == expected + assert config.max_per_dataset == 5 + assert original.max_total == 20 + + def test_estimate_cache_normalizes_default_limits_but_preserves_all(self) -> None: + scenario_class = MagicMock() + keys = { + ScenarioService._build_configured_estimate_key( + scenario_name="test", scenario_class=scenario_class, request=ScenarioRunSizeEstimateRequest(**args) + ) + for args in ( + {}, + {"max_dataset_size": None}, + {"max_dataset_size": ""}, + {"max_dataset_size": "all"}, + {"max_dataset_size": 7}, + ) + } + assert len(keys) == 3 + async def test_configured_estimate_uses_shared_launch_resolution(self) -> None: """Configured estimates pass typed selections and parameters into the registry lifecycle.""" metadata = _make_scenario_metadata(registry_name="airt.jailbreak") diff --git a/tests/unit/cli/test_api_client.py b/tests/unit/cli/test_api_client.py index 3b5b7244eb..7185fbf16d 100644 --- a/tests/unit/cli/test_api_client.py +++ b/tests/unit/cli/test_api_client.py @@ -733,13 +733,35 @@ async def test_start_scenario_run_async(client, mock_httpx_client): mock_httpx_client.post.assert_awaited_once() args, kwargs = mock_httpx_client.post.call_args assert args == ("/api/scenarios/runs",) - # The CLI serializes the typed request via model_dump(mode="json", exclude_none=True); - # required fields must appear in the body, None-valued fields must not. assert kwargs["json"]["scenario_name"] == "x" assert kwargs["json"]["target_name"] == "t" assert "scenario_params" not in kwargs["json"] +@pytest.mark.parametrize( + "limit_args", + [ + {}, + {"max_dataset_size": None}, + {"max_dataset_size": ""}, + {"max_dataset_size": "default"}, + {"max_dataset_size": "all"}, + {"max_dataset_size": 7}, + ], +) +async def test_start_run_preserves_explicit_all_async( + *, client: PyRITApiClient, mock_httpx_client: MagicMock, limit_args: dict[str, int | str | None] +) -> None: + mock_httpx_client.post.return_value = _make_response(json_data=_run_summary_payload()) + await client.start_scenario_run_async( + request=RunScenarioRequest(scenario_name="test", target_name="target", **limit_args) + ) + payload = mock_httpx_client.post.call_args.kwargs["json"] + expected = limit_args.get("max_dataset_size") or "default" + assert "max_dataset_size" in payload + assert payload.get("max_dataset_size") == expected + + async def test_get_scenario_run_async(client, mock_httpx_client): import httpx as _httpx diff --git a/tests/unit/cli/test_cli_args.py b/tests/unit/cli/test_cli_args.py index 1f7ca53e74..749f0c0a28 100644 --- a/tests/unit/cli/test_cli_args.py +++ b/tests/unit/cli/test_cli_args.py @@ -3,10 +3,23 @@ import pytest -from pyrit.cli._cli_args import _argparse_validator, parse_run_arguments +from pyrit.cli._cli_args import _argparse_validator, parse_list_targets_arguments, parse_run_arguments from pyrit.models import Parameter +def test_shell_parsers_omit_unsupplied_arguments() -> None: + assert parse_run_arguments(args_string="test") == {"scenario_name": "test"} + assert parse_list_targets_arguments(args_string="") == {} + + +@pytest.mark.parametrize(("value", "expected"), [("all", "all"), ("default", "default"), ("7", 7), ('""', "default")]) +def test_shell_explicit_dataset_limit(*, value: str, expected: int | str | None) -> None: + assert parse_run_arguments(args_string=f"test --max-dataset-size {value}") == { + "scenario_name": "test", + "max_dataset_size": expected, + } + + def _sp(*, name: str, description: str = "", param_type: str = "str") -> Parameter: """Build a real Parameter from Summary-style kwargs (param_type as a string).""" return Parameter.model_validate( diff --git a/tests/unit/cli/test_pyrit_scan.py b/tests/unit/cli/test_pyrit_scan.py index 86360f52a7..e9fa072d37 100644 --- a/tests/unit/cli/test_pyrit_scan.py +++ b/tests/unit/cli/test_pyrit_scan.py @@ -14,6 +14,7 @@ from pyrit.cli import _config_reader as pyrit_scan_config_reader from pyrit.cli import pyrit_scan +from pyrit.cli._cli_args import parse_run_arguments from pyrit.models import Parameter from unit.mocks import make_scenario_result @@ -804,6 +805,19 @@ def test_skips_params_that_collide_with_existing_flags(self): class TestBuildRunRequest: """Tests for _build_run_request.""" + @pytest.mark.parametrize("value", [None, "default", "all", "7", ""]) + def test_limit_presence_matches_shell_parser(self, value: str | None) -> None: + flags = [] if value is None else ["--max-dataset-size", value] + parsed = pyrit_scan.parse_args(["run", "test", "--target", "target", *flags]) + request = pyrit_scan._build_run_request(parsed_args=parsed, scenario_name="test") + shell_args = parse_run_arguments( + args_string="test --target target " + ("--max-dataset-size ''" if value == "" else " ".join(flags)) + ) + assert ("max_dataset_size" in request.model_fields_set) == (value is not None) + assert ("max_dataset_size" in shell_args) == (value is not None) + expected = 7 if value == "7" else value or "default" + assert request.max_dataset_size == shell_args.get("max_dataset_size", "default") == expected + def test_includes_initializer_args(self): parsed = Namespace( target="t", diff --git a/tests/unit/cli/test_pyrit_shell.py b/tests/unit/cli/test_pyrit_shell.py index 77f5dda66e..d678eaa8f3 100644 --- a/tests/unit/cli/test_pyrit_shell.py +++ b/tests/unit/cli/test_pyrit_shell.py @@ -1190,6 +1190,34 @@ def test_stop_server_close_client_swallows_errors(self, shell): class TestShellScenarioParamFlow: """Regression tests: shell.do_run must forward scenario-declared parameters.""" + @pytest.mark.parametrize( + "flag", + [ + "", + " --max-dataset-size default", + " --max-dataset-size all", + " --max-dataset-size 7", + " --max-dataset-size ''", + ], + ) + def test_run_preserves_total_limit_presence( + self, *, shell: tuple[pyrit_shell.PyRITShell, AsyncMock], flag: str + ) -> None: + s, client = shell + client.start_scenario_run_async = AsyncMock(return_value=TestDoRun._run_payload("CREATED")) + client.get_scenario_run_async = AsyncMock(return_value=TestDoRun._run_payload("COMPLETED")) + client.get_scenario_run_results_async = AsyncMock(return_value=TestDoRun._empty_scenario_result()) + with ( + patch("pyrit.cli._output.print_scenario_result_async", new_callable=AsyncMock), + patch("pyrit.cli._output.print_scenario_run_progress"), + patch("time.sleep"), + ): + s.do_run(f"foo --target t{flag}") + request = client.start_scenario_run_async.call_args.kwargs["request"] + expected = 7 if flag.endswith("7") else "all" if flag.endswith("all") else "default" + assert ("max_dataset_size" in request.model_fields_set) == bool(flag) + assert request.max_dataset_size == expected + def test_run_passes_scenario_declared_params(self, shell): s, client = shell client.get_scenario_async.return_value = client._make_typed_scenario( diff --git a/tests/unit/datasets/test_seed_dataset_provider.py b/tests/unit/datasets/test_seed_dataset_provider.py index bcbd113584..973019d72c 100644 --- a/tests/unit/datasets/test_seed_dataset_provider.py +++ b/tests/unit/datasets/test_seed_dataset_provider.py @@ -64,6 +64,20 @@ def mock_darkbench_data(): class TestSeedDatasetProvider: """Test the SeedDatasetProvider base class and registration.""" + async def test_name_lookup_does_not_parse_metadata_or_fetch(self) -> None: + factory = MagicMock() + factory.return_value.dataset_name = "known" + factory.return_value._parse_metadata_async = AsyncMock(side_effect=AssertionError("metadata read")) + factory.return_value.fetch_dataset_async = AsyncMock(side_effect=AssertionError("fetch")) + with ( + patch.dict(SeedDatasetProvider._registry, {"test": factory}, clear=True), + patch.object(SeedDatasetProvider, "_materialize_builtin_providers"), + ): + result = await SeedDatasetProvider.get_providers_by_name_async(dataset_names=["known", "missing"]) + assert result == {"known": factory.return_value} + factory.return_value._parse_metadata_async.assert_not_awaited() + factory.return_value.fetch_dataset_async.assert_not_awaited() + def test_registration(self): """Test that subclasses are automatically registered.""" diff --git a/tests/unit/memory/memory_interface/test_interface_seed_prompts.py b/tests/unit/memory/memory_interface/test_interface_seed_prompts.py index 79dd24668d..49738452ad 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_prompts.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_prompts.py @@ -8,7 +8,7 @@ from uuid import uuid4 import pytest -from sqlalchemy import String +from sqlalchemy import String, event from sqlalchemy.exc import SQLAlchemyError from pyrit.memory import MemoryInterface @@ -818,6 +818,25 @@ async def test_get_seed_dataset_names_multiple(sqlite_instance: MemoryInterface) assert sorted(await sqlite_instance.get_seed_dataset_names_async()) == sorted(dataset_names) +async def test_get_seed_dataset_names_without_loading_seed_rows_async(sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[ + SeedObjective(value=f"objective-{index}", dataset_name=name) + for index, name in enumerate(["first", "first", "second", None, ""]) + ], + added_by="test", + ) + + def reject_seed_load(*_: object) -> None: + raise AssertionError("Listing dataset names must not load full seed rows") + + event.listen(SeedEntry, "load", reject_seed_load) + try: + assert set(await sqlite_instance.get_seed_dataset_names_async()) == {"first", "second"} + finally: + event.remove(SeedEntry, "load", reject_seed_load) + + async def test_add_seed_groups_to_memory_empty_list(sqlite_instance: MemoryInterface): prompt_group = SeedGroup(seeds=[SeedPrompt(value="Test prompt", added_by="tester", data_type="text", sequence=0)]) prompt_group.seeds = [] diff --git a/tests/unit/models/test_scenario_request.py b/tests/unit/models/test_scenario_request.py index 938ed2ef04..b724294b5d 100644 --- a/tests/unit/models/test_scenario_request.py +++ b/tests/unit/models/test_scenario_request.py @@ -6,7 +6,48 @@ import pytest from pydantic import ValidationError -from pyrit.models.catalog.scenario import DATASET_FILTERS, RunScenarioRequest +from pyrit.models.catalog.scenario import DATASET_FILTERS, RunScenarioRequest, ScenarioRunSizeEstimateRequest + + +@pytest.mark.parametrize("model", [RunScenarioRequest, ScenarioRunSizeEstimateRequest]) +@pytest.mark.parametrize( + ("limit", "expected"), + [ + (None, "default"), + ("", "default"), + (" ", "default"), + ("default", "default"), + (" DEFAULT ", "default"), + ("all", "all"), + (" ALL ", "all"), + (7, 7), + ("7", 7), + ], +) +def test_dataset_limit_values( + *, model: type[RunScenarioRequest] | type[ScenarioRunSizeEstimateRequest], limit: object, expected: object +) -> None: + request = model.model_validate({"scenario_name": "s", "target_name": "t", "max_dataset_size": limit}) + assert request.max_dataset_size == expected + assert request.model_dump(mode="json")["max_dataset_size"] == expected + + +@pytest.mark.parametrize("model", [RunScenarioRequest, ScenarioRunSizeEstimateRequest]) +@pytest.mark.parametrize("limit", [0, -1, True, False, 1.0, 1.5, "unlimited", [], {}]) +def test_dataset_limit_rejects_invalid_values( + *, model: type[RunScenarioRequest] | type[ScenarioRunSizeEstimateRequest], limit: object +) -> None: + with pytest.raises(ValidationError): + model.model_validate({"scenario_name": "s", "target_name": "t", "max_dataset_size": limit}) + + +@pytest.mark.parametrize("model", [RunScenarioRequest, ScenarioRunSizeEstimateRequest]) +def test_omitted_limit_uses_explicit_default( + model: type[RunScenarioRequest] | type[ScenarioRunSizeEstimateRequest], +) -> None: + request = model.model_validate({"scenario_name": "s", "target_name": "t"}) + assert request.max_dataset_size == "default" + assert request.model_dump(mode="json")["max_dataset_size"] == "default" def _make_request(*, dataset_filters: dict[str, list[str]] | None) -> RunScenarioRequest: diff --git a/tests/unit/scenario/airt/test_cyber.py b/tests/unit/scenario/airt/test_cyber.py index 7a94fa653c..3521d3e1f6 100644 --- a/tests/unit/scenario/airt/test_cyber.py +++ b/tests/unit/scenario/airt/test_cyber.py @@ -217,7 +217,7 @@ async def test_initialize_raises_when_no_datasets(self, mock_objective_target, m # Neutralize the provider fetch so the empty-memory path raises loudly instead of fetching # the real default dataset from the provider. with patch( - "pyrit.scenario.core.dataset_configuration.DatasetConfiguration._fetch_dataset_async", + "pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async", new_callable=AsyncMock, ): scenario.set_params_from_args(args={"objective_target": mock_objective_target}) diff --git a/tests/unit/scenario/airt/test_jailbreak.py b/tests/unit/scenario/airt/test_jailbreak.py index 929cd0867b..7543223d70 100644 --- a/tests/unit/scenario/airt/test_jailbreak.py +++ b/tests/unit/scenario/airt/test_jailbreak.py @@ -232,7 +232,9 @@ async def test_run_size_prompt_sending_two_templates_four_groups_is_eight( assert [component.label for component in estimate.components] == ["Inline jailbreak delivery"] assert estimate.datasets[0].logical_seed_group_count is None assert estimate.datasets[0].selected_seed_group_count is None - assert [(cap.label, cap.count) for cap in estimate.datasets[0].configured_caps] == [("per-dataset cap", 4)] + assert [(cap.label, cap.count) for cap in estimate.datasets[0].configured_caps] == [ + ("combined configuration cap", 4) + ] async def test_run_size_is_conditional_when_system_delivery_target_is_not_selected( self, mock_objective_scorer @@ -356,7 +358,7 @@ async def test_init_raises_exception_when_no_datasets_available(self, mock_objec scenario = Jailbreak(objective_scorer=mock_objective_scorer) with patch( - "pyrit.scenario.core.dataset_configuration.DatasetConfiguration._fetch_dataset_async", + "pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async", new_callable=AsyncMock, ): scenario.set_params_from_args(args={"objective_target": mock_objective_target}) diff --git a/tests/unit/scenario/airt/test_leakage.py b/tests/unit/scenario/airt/test_leakage.py index e4e9f3cb20..1888d3be4f 100644 --- a/tests/unit/scenario/airt/test_leakage.py +++ b/tests/unit/scenario/airt/test_leakage.py @@ -49,6 +49,9 @@ def mock_dataset_config(mock_memory_seeds): """Create a mock dataset config that returns the seed groups.""" seed_groups = [AttackSeedGroup(seeds=[seed]) for seed in mock_memory_seeds] mock_config = MagicMock(spec=DatasetAttackConfiguration) + mock_config.with_overrides.return_value = mock_config + mock_config.sources = () + mock_config.max_total = "all" mock_config.get_attack_groups_by_dataset_async = AsyncMock(return_value={"airt_leakage": seed_groups}) mock_config.dataset_names = ["airt_leakage"] return mock_config diff --git a/tests/unit/scenario/airt/test_psychosocial.py b/tests/unit/scenario/airt/test_psychosocial.py index 1bb1f2c438..f8d720d373 100644 --- a/tests/unit/scenario/airt/test_psychosocial.py +++ b/tests/unit/scenario/airt/test_psychosocial.py @@ -27,6 +27,8 @@ from pyrit.scenario.core.dataset_configuration import ( CompoundDatasetAttackConfiguration, DatasetAttackConfiguration, + DatasetConstraintError, + DatasetSource, ) from pyrit.scenario.core.scenario import Scenario from pyrit.scenario.scenarios.airt.psychosocial import ( @@ -121,11 +123,52 @@ def register_default_targets(): FIXTURES = ["patch_central_database"] +@pytest.mark.usefixtures(*FIXTURES) +@pytest.mark.parametrize("total, expected", [(None, 8), (3, 3)]) +async def test_unequal_source_caps_and_total_estimate(total: int | None, expected: int) -> None: + scenario = _scenario_with_mock_scorers() + scenario.set_params_from_args( + args={ + "dataset_config": DatasetAttackConfiguration( + sources=[ + DatasetSource(name=_SUB_HARMS[0].dataset_name, max_size=1), + DatasetSource(name=_SUB_HARMS[1].dataset_name, max_size=7), + ], + max_total=total, + ), + "scenario_techniques": [PsychosocialTechnique.NoConverter], + "include_baseline": True, + } + ) + with ( + patch.object(CentralMemory.get_memory_instance(), "get_seeds_async", side_effect=AssertionError("No reads")), + patch.object(DatasetAttackConfiguration, "prepare_async", side_effect=AssertionError("No preparation")), + ): + estimate = await scenario.get_run_size_estimate_async() + assert estimate.estimated_attack_count == expected * 2 + + +@pytest.mark.usefixtures(*FIXTURES) +async def test_compound_sources_require_explicit_migration() -> None: + scenario = _scenario_with_mock_scorers() + scenario.set_params_from_args( + args={ + "dataset_config": CompoundDatasetAttackConfiguration( + configurations=[ + DatasetAttackConfiguration(sources=[DatasetSource(name=harm.dataset_name)]) for harm in _SUB_HARMS + ] + ), + } + ) + with pytest.raises(DatasetConstraintError, match="compound children"): + await scenario.get_run_size_estimate_async() + + @pytest.mark.usefixtures(*FIXTURES) @pytest.mark.parametrize("outer_limit", [3, None]) @pytest.mark.parametrize("sub_harm", ["all", "imminent_crisis"]) @pytest.mark.parametrize("include_baseline", [False, True]) -async def test_compound_estimate_uses_initialization_cap_async( +async def test_total_estimate_uses_initialization_cap_async( *, mock_objective_target: PromptTarget, outer_limit: int | None, @@ -142,11 +185,11 @@ async def test_compound_estimate_uses_initialization_cap_async( def get_seeds(*, dataset_name: str, **_: object) -> list[SeedObjective]: return list(seeds_by_dataset[dataset_name]) - config = CompoundDatasetAttackConfiguration.per_dataset( - dataset_names=list(seeds_by_dataset), - max_dataset_size=1 if outer_limit is not None else 2, + config = DatasetAttackConfiguration( + sources=[DatasetSource(name=name) for name in seeds_by_dataset], + max_per_dataset="all", + max_total=outer_limit, ) - config.max_dataset_size = outer_limit scenario = _scenario_with_mock_scorers() scenario.set_params_from_args( args={ @@ -158,8 +201,8 @@ def get_seeds(*, dataset_name: str, **_: object) -> list[SeedObjective]: } ) harm_count = 2 if sub_harm == "all" else 1 - per_harm_count = outer_limit if outer_limit is not None else 6 - expected_count = per_harm_count * harm_count * (1 + include_baseline) + selected_count = outer_limit if outer_limit is not None else 6 * harm_count + expected_count = selected_count * (1 + include_baseline) memory = CentralMemory.get_memory_instance() with patch.object( memory, @@ -173,20 +216,23 @@ def get_seeds(*, dataset_name: str, **_: object) -> list[SeedObjective]: if outer_limit is None: assert estimate.status is ScenarioRunSizeEstimateStatus.Unavailable assert estimate.estimated_attack_count is None - assert estimate.dataset_size == scenario_dataset_size_from_limit(None) + assert estimate.dataset_size == scenario_dataset_size_from_limit("all") assert all(dataset.configured_caps == [] for dataset in estimate.datasets) else: - assert estimate.dataset_size == scenario_dataset_size_from_limit(outer_limit * harm_count) + assert estimate.dataset_size == scenario_dataset_size_from_limit(outer_limit) assert estimate.dataset_limit.value == outer_limit assert [dataset.configured_caps[0].count for dataset in estimate.datasets] == [outer_limit] * harm_count assert estimate.status is ScenarioRunSizeEstimateStatus.Approximate assert estimate.estimated_attack_count == expected_count - with patch.object(memory, "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds) as read_seeds: + with ( + patch.object(memory, "get_seed_dataset_names_async", return_value=list(seeds_by_dataset)), + patch.object(memory, "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds) as read_seeds, + ): await scenario.initialize_async() read_seeds.assert_awaited() plan = scenario._build_run_plan() - assert len(plan.seed_groups) == per_harm_count * harm_count + assert len(plan.seed_groups) == selected_count assert sum(len(group.seed_group_ids) for group in plan.atomic_groups) == expected_count @@ -333,7 +379,7 @@ def test_no_arg_construct_works(self): assert Psychosocial().uses_default_adversarial_target is True def test_version_is_3(self): - assert Psychosocial.VERSION == 4 + assert Psychosocial.VERSION == 5 def test_default_technique_is_default(self): assert _scenario_with_mock_scorers()._default_technique == PsychosocialTechnique.DEFAULT @@ -493,9 +539,9 @@ async def test_single_sub_harm_hard_binds_dataset_config(self, mock_objective_ta await scenario.initialize_async() assert list(scenario._dataset_config.dataset_names) == ["airt_licensed_therapist"] - async def test_max_dataset_size_applied_per_sub_harm_when_hard_binding(self, mock_objective_target): + async def test_unrelated_source_is_rejected(self, mock_objective_target): scenario = _scenario_with_mock_scorers() - with _patch_base_seed_groups(_make_seed_groups()): + with pytest.raises(DatasetConstraintError, match="sub-harm datasets"): scenario.set_params_from_args( args={ "objective_target": mock_objective_target, @@ -503,18 +549,9 @@ async def test_max_dataset_size_applied_per_sub_harm_when_hard_binding(self, moc } ) await scenario.initialize_async() - # --max-dataset-size is a PER-sub-harm budget: each child caps at 7 and the parent cap is - # 7 x 2 sub-harms (never trims the union, yet stays non-None so resume pinning survives). - assert isinstance(scenario._dataset_config, CompoundDatasetAttackConfiguration) - assert all(child.max_dataset_size == 7 for child in scenario._dataset_config._configurations) - assert scenario._dataset_config.max_dataset_size == 14 - assert set(scenario._dataset_config.dataset_names) == { - "airt_imminent_crisis", - "airt_licensed_therapist", - } - async def test_max_dataset_size_one_keeps_both_sub_harms(self, mock_objective_target): - """A global budget of 1 starves a sub-harm; the per-sub-harm compound keeps both. + async def test_per_dataset_limit_one_keeps_both_sub_harms(self, mock_objective_target): + """A cap of one per source keeps both sub-harms. Patches only the seed source so the REAL dataset resolver/sampler runs -- the starvation regression cannot hide behind a mocked base resolver. @@ -535,24 +572,26 @@ def _get_seeds(*, dataset_name, **_): memory = CentralMemory.get_memory_instance() scenario = _scenario_with_mock_scorers() - with patch.object(memory, "get_seeds_async", side_effect=_get_seeds): + with ( + patch.object(memory, "get_seed_dataset_names_async", return_value=list(seeds_by_dataset)), + patch.object(memory, "get_seeds_async", side_effect=_get_seeds), + ): scenario.set_params_from_args( args={ "objective_target": mock_objective_target, "sub_harm": "all", "scenario_techniques": [PsychosocialTechnique.NoConverter], "dataset_config": DatasetAttackConfiguration( - dataset_names=["ignored"], max_dataset_size=1, auto_fetch=False + sources=[DatasetSource(name=harm.dataset_name) for harm in _SUB_HARMS], + max_per_dataset=1, + auto_fetch=False, ), } ) await scenario.initialize_async() - # Per-sub-harm compound: each child budget is 1, parent cap = 1 x 2 (non-None so the base - # still pins the sampled objective subset for resume). - assert isinstance(scenario._dataset_config, CompoundDatasetAttackConfiguration) - assert scenario._dataset_config.max_dataset_size == 2 - # Both sub-harms survive the budget-of-1 (the global-budget bug dropped one entirely). + assert scenario._dataset_config.max_per_dataset == 1 + assert scenario._dataset_config.max_total == "all" assert {a.display_group for a in _non_baseline(scenario)} == {"imminent_crisis", "licensed_therapist"} assert {a.atomic_attack_name for a in _baselines(scenario)} == { "imminent_crisis_baseline", @@ -575,16 +614,21 @@ def get_seeds(*, dataset_name: str, **_: object) -> list[SeedObjective]: args={ "objective_target": mock_objective_target, "scenario_techniques": [PsychosocialTechnique.NoConverter], - "dataset_config": DatasetAttackConfiguration(dataset_names=["ignored"], max_dataset_size=None), + "dataset_config": DatasetAttackConfiguration( + sources=[DatasetSource(name=harm.dataset_name) for harm in _SUB_HARMS], + max_per_dataset="all", + ), } ) - with patch.object( - CentralMemory.get_memory_instance(), "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds - ) as read_seeds: + memory = CentralMemory.get_memory_instance() + with ( + patch.object(memory, "get_seed_dataset_names_async", return_value=list(seeds_by_dataset)), + patch.object(memory, "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds) as read_seeds, + ): await scenario.initialize_async() read_seeds.assert_awaited() - assert scenario._dataset_config.max_dataset_size is None + assert scenario._dataset_config.max_dataset_size == "all" assert len(scenario._atomic_attacks) == 4 assert all(len(attack.seed_groups) == 6 for attack in scenario._atomic_attacks) plan = scenario._build_run_plan() diff --git a/tests/unit/scenario/airt/test_rapid_response.py b/tests/unit/scenario/airt/test_rapid_response.py index 1f40f58234..072abb793a 100644 --- a/tests/unit/scenario/airt/test_rapid_response.py +++ b/tests/unit/scenario/airt/test_rapid_response.py @@ -19,7 +19,7 @@ from pyrit.registry import TargetRegistry from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration from pyrit.scenario.scenarios.airt.rapid_response import RapidResponse from pyrit.score import TrueFalseScorer from pyrit.setup.initializers.techniques import ( @@ -185,7 +185,7 @@ def test_default_dataset_config_has_all_harm_datasets(self, mock_objective_score "pyrit.scenario.core.scenario.Scenario._get_default_objective_scorer", return_value=mock_objective_scorer ): config = RapidResponse()._default_dataset_config - assert isinstance(config, CompoundDatasetAttackConfiguration) + assert isinstance(config, DatasetAttackConfiguration) names = config.dataset_names expected = [f"airt_{cat}" for cat in ALL_HARM_CATEGORIES] for name in expected: @@ -197,7 +197,8 @@ def test_default_dataset_config_max_dataset_size(self, mock_objective_scorer): "pyrit.scenario.core.scenario.Scenario._get_default_objective_scorer", return_value=mock_objective_scorer ): config = RapidResponse()._default_dataset_config - assert all(child.max_dataset_size == 4 for child in config._configurations) + assert config.max_per_dataset == 4 + assert config.max_total == "all" @patch("pyrit.scenario.core.scenario.Scenario._get_default_objective_scorer") def test_initialization_minimal(self, mock_get_scorer, mock_objective_scorer): @@ -213,7 +214,7 @@ def test_initialization_with_custom_scorer(self, mock_objective_scorer): @patch("pyrit.scenario.core.scenario.Scenario._get_default_objective_scorer") @patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=ALL_HARM_SEED_GROUPS, @@ -243,7 +244,7 @@ async def test_initialize_raises_when_no_datasets(self, mock_objective_target, m # Neutralize the provider fetch so the empty-memory path raises loudly instead of fetching # the real default dataset from the provider. with patch( - "pyrit.scenario.core.dataset_configuration.DatasetConfiguration._fetch_dataset_async", + "pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async", new_callable=AsyncMock, ): scenario.set_params_from_args(args={"objective_target": mock_objective_target}) @@ -252,7 +253,7 @@ async def test_initialize_raises_when_no_datasets(self, mock_objective_target, m @patch("pyrit.scenario.core.scenario.Scenario._get_default_objective_scorer") @patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=ALL_HARM_SEED_GROUPS, @@ -302,7 +303,7 @@ async def _init_and_get_attacks( """Helper: initialize scenario and return atomic attacks.""" groups = seed_groups or {"hate": _make_seed_groups("hate")} with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -419,7 +420,7 @@ def _spy_create(self, **kwargs): groups = {"hate": _make_seed_groups("hate")} with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -518,7 +519,7 @@ async def test_unknown_technique_skipped_with_warning(self, mock_objective_targe ) with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, diff --git a/tests/unit/scenario/airt/test_scam.py b/tests/unit/scenario/airt/test_scam.py index 6ff188c317..5c4a99043e 100644 --- a/tests/unit/scenario/airt/test_scam.py +++ b/tests/unit/scenario/airt/test_scam.py @@ -54,6 +54,9 @@ def mock_dataset_config(mock_memory_seed_groups): """Create a mock dataset config that returns the seed groups.""" attack_seed_groups = list(mock_memory_seed_groups) mock_config = MagicMock(spec=DatasetAttackConfiguration) + mock_config.with_overrides.return_value = mock_config + mock_config.sources = () + mock_config.max_total = "all" mock_config.get_attack_seed_groups_async = AsyncMock(return_value=attack_seed_groups) mock_config.get_attack_groups_by_dataset_async = AsyncMock(return_value={"airt_scam": attack_seed_groups}) mock_config.dataset_names = ["airt_scam"] @@ -231,7 +234,7 @@ async def test_init_raises_exception_when_no_datasets_available_async( # Error should occur during initialize_async when _get_atomic_attacks_async resolves seed groups. # Neutralize the provider fetch so the empty-memory path raises loudly instead of fetching. with patch( - "pyrit.scenario.core.dataset_configuration.DatasetConfiguration._fetch_dataset_async", + "pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async", new_callable=AsyncMock, ): scenario.set_params_from_args(args={"objective_target": mock_objective_target}) diff --git a/tests/unit/scenario/benchmark/test_adversarial.py b/tests/unit/scenario/benchmark/test_adversarial.py index 18539aafdd..fe08b97942 100644 --- a/tests/unit/scenario/benchmark/test_adversarial.py +++ b/tests/unit/scenario/benchmark/test_adversarial.py @@ -75,7 +75,13 @@ from pyrit.prompt_target import PromptTarget from pyrit.registry import TargetRegistry from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry -from pyrit.scenario.core import AtomicAttack, BaselineAttackPolicy, CompoundDatasetAttackConfiguration +from pyrit.scenario.core import ( + AtomicAttack, + BaselineAttackPolicy, + CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, + DatasetSource, +) from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory from pyrit.scenario.core.scenario import Scenario from pyrit.scenario.scenarios.benchmark.adversarial import ( @@ -118,6 +124,37 @@ def _build_benchmarkable_factories_snapshot() -> list: "tap", } + +@pytest.mark.usefixtures("patch_central_database") +async def test_named_source_limits_keep_benchmark_selection_stable() -> None: + bench = AdversarialBenchmark(objective_scorer=MagicMock(spec=TrueFalseScorer)) + bench._dataset_config = DatasetAttackConfiguration( + sources=[DatasetSource(name="a", max_size=2), DatasetSource(name="b", max_size=3)], + max_total=4, + ) + groups = { + name: [ + AttackSeedGroup(seeds=[SeedObjective(value=f"{name}-{index}", harm_categories=[str(index % 2)])]) + for index in range(8) + ] + for name in ("a", "b") + } + with patch.object( + DatasetAttackConfiguration, + "get_attack_groups_by_dataset_async", + side_effect=lambda **_: {name: list(items) for name, items in groups.items()}, + ) as read: + first = await bench._resolve_seed_groups_by_dataset_async() + second = await bench._resolve_seed_groups_by_dataset_async() + full = await bench._resolve_seed_groups_by_dataset_async(apply_sampling=False) + assert first == second + assert sum(map(len, first.values())) == 4 + assert len(first.get("a", [])) <= 2 + assert len(first.get("b", [])) <= 3 + assert sum(map(len, full.values())) == 16 + assert all(call.kwargs == {"apply_sampling": False} for call in read.await_args_list) + + # --------------------------------------------------------------------------- # Fixtures / helpers # --------------------------------------------------------------------------- @@ -232,12 +269,15 @@ def get_seeds(*, dataset_name: str, **_: object) -> list[SeedObjective]: assert estimate.estimated_attack_count == outer_limit expected_count = outer_limit if outer_limit is not None else 6 - assert estimate.dataset_size == scenario_dataset_size_from_limit(outer_limit) + assert estimate.dataset_size == scenario_dataset_size_from_limit(outer_limit if outer_limit is not None else "all") assert all( [cap.count for cap in dataset.configured_caps] == ([] if outer_limit is None else [outer_limit]) for dataset in estimate.datasets ) - with patch.object(bench._memory, "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds) as read_seeds: + with ( + patch.object(bench._memory, "get_seed_dataset_names_async", return_value=list(seeds_by_dataset)), + patch.object(bench._memory, "get_seeds_async", new_callable=AsyncMock, side_effect=get_seeds) as read_seeds, + ): await bench.initialize_async() read_seeds.assert_awaited() plan = bench._build_run_plan() @@ -575,6 +615,7 @@ async def test_initialize_without_selection_resolves_exact_default(self): ) with ( + patch.object(DatasetAttackConfiguration, "prepare_async", new_callable=AsyncMock), patch.object(bench, "_resolve_seed_groups_by_dataset_async", new_callable=AsyncMock, return_value={}), patch.object(bench, "_build_atomic_attacks_async", new_callable=AsyncMock, return_value=[]), ): @@ -916,8 +957,9 @@ def _make_bench_with_targets( # Dataset config: one dataset with one real seed group (AtomicAttack hashes objectives). seed_group = AttackSeedGroup(seeds=[SeedObjective(value="benchmark_objective_1")]) - bench._dataset_config = MagicMock() - bench._dataset_config.max_dataset_size = None + bench._dataset_config = MagicMock(spec=DatasetAttackConfiguration) + bench._dataset_config.sources = () + bench._dataset_config.max_total = "all" bench._dataset_config.get_attack_groups_by_dataset_async = AsyncMock(return_value={"harmbench": [seed_group]}) return bench @@ -964,8 +1006,9 @@ async def test_display_group_uses_registry_name_not_target_model_name(self): bench._scenario_techniques = [red_teaming_technique] seed_group = AttackSeedGroup(seeds=[SeedObjective(value="display_group_regression_objective")]) - bench._dataset_config = MagicMock() - bench._dataset_config.max_dataset_size = None + bench._dataset_config = MagicMock(spec=DatasetAttackConfiguration) + bench._dataset_config.sources = () + bench._dataset_config.max_total = "all" bench._dataset_config.get_attack_groups_by_dataset_async = AsyncMock(return_value={"harmbench": [seed_group]}) result = await _build_atomic_attacks(bench) @@ -1402,8 +1445,9 @@ def _make_bench(self, *, use_cached: bool) -> AdversarialBenchmark: bench._scenario_techniques = [red_teaming_technique] seed_group = AttackSeedGroup(seeds=[SeedObjective(value="skip_cached_objective")]) - bench._dataset_config = MagicMock() - bench._dataset_config.max_dataset_size = None + bench._dataset_config = MagicMock(spec=DatasetAttackConfiguration) + bench._dataset_config.sources = () + bench._dataset_config.max_total = "all" bench._dataset_config.get_attack_groups_by_dataset_async = AsyncMock(return_value={"harmbench": [seed_group]}) return bench @@ -1415,7 +1459,7 @@ def _patch_identifier(self, eval_hash: str = "obj_hash"): async def test_sampling_is_stable_across_fresh_runs(self): bench = self._make_bench(use_cached=False) - bench._dataset_config.max_dataset_size = 1 + bench._dataset_config.max_total = 1 group_a = AttackSeedGroup(seeds=[SeedObjective(value="objective a")]) group_b = AttackSeedGroup(seeds=[SeedObjective(value="objective b")]) @@ -1436,7 +1480,7 @@ async def test_sampling_is_stable_across_fresh_runs(self): async def test_sampling_balances_single_harm_categories_without_cache(self) -> None: bench = self._make_bench(use_cached=False) - bench._dataset_config.max_dataset_size = 24 + bench._dataset_config.max_total = 24 categories = [ "election_critical_information", "hate_v3", diff --git a/tests/unit/scenario/core/test_dataset_configuration.py b/tests/unit/scenario/core/test_dataset_configuration.py index 187d2ba4cd..7d7311009c 100644 --- a/tests/unit/scenario/core/test_dataset_configuration.py +++ b/tests/unit/scenario/core/test_dataset_configuration.py @@ -11,6 +11,7 @@ from pyrit.models import ( AttackSeedGroup, IndeterminateDatasetSize, + SeedDataset, SeedGroup, SeedObjective, SeedPrompt, @@ -22,10 +23,10 @@ DatasetAttackConfiguration, DatasetConfiguration, DatasetConstraintError, + DatasetFetchPolicy, DatasetSourceKind, ResolvedDataset, forbid_inline_seeds, - read_only_dataset_resolution, require_harm_categories, require_inline_seeds, require_min_size, @@ -51,6 +52,7 @@ def resolved( def mock_memory() -> MagicMock: """A stand-in CentralMemory whose ``get_seeds`` returns nothing by default.""" memory = MagicMock(spec=MemoryInterface) + memory.get_seed_dataset_names_async = AsyncMock(return_value=[]) memory.get_seeds_async = AsyncMock(return_value=[]) memory.get_seed_groups_async = AsyncMock(return_value=[]) memory.add_seed_datasets_to_memory_async = AsyncMock() @@ -86,9 +88,10 @@ def make_objectives(*values: str) -> list[SeedObjective]: @pytest.mark.parametrize( - ("kwargs", "expected"), [({}, 5), ({"max_dataset_size": None}, 8), ({"max_dataset_size": 2}, 2)] + ("kwargs", "expected"), + [({}, 5), ({"max_dataset_size": None}, 5), ({"max_dataset_size": "all"}, 8), ({"max_dataset_size": 2}, 2)], ) -async def test_default_limit_and_explicit_overrides(*, kwargs: dict[str, int | None], expected: int) -> None: +async def test_default_limit_and_explicit_overrides(*, kwargs: dict[str, int | str | None], expected: int) -> None: config = DatasetAttackConfiguration(seeds=make_objectives(*(str(index) for index in range(8))), **kwargs) groups = await config.get_attack_seed_groups_async() assert len(groups) == expected @@ -107,11 +110,11 @@ def test_compound_budget_combines_children_before_outer_cap(*, outer_limit: int assert config.get_size_budget() == scenario_dataset_size_from_limit(expected) -@pytest.mark.parametrize(("outer_limit", "expected"), [(None, None), (7, 7)]) -def test_unlimited_child_budget_needs_outer_limit(*, outer_limit: int | None, expected: int | None) -> None: +@pytest.mark.parametrize(("outer_limit", "expected"), [(None, "all"), (7, 7)]) +def test_unlimited_child_budget_needs_outer_limit(*, outer_limit: int | None, expected: int | str) -> None: config = CompoundDatasetAttackConfiguration( configurations=[ - DatasetAttackConfiguration(dataset_names=["a"], max_dataset_size=None), + DatasetAttackConfiguration(dataset_names=["a"], max_per_dataset="all", max_total="all"), DatasetAttackConfiguration(dataset_names=["b"]), ], max_dataset_size=outer_limit, @@ -122,13 +125,13 @@ def test_unlimited_child_budget_needs_outer_limit(*, outer_limit: int | None, ex def test_per_dataset_default_does_not_add_implicit_compound_cap() -> None: config = CompoundDatasetAttackConfiguration.per_dataset(dataset_names=["a", "b"]) assert config.get_size_budget() == scenario_dataset_size_from_limit(10) - assert config.max_dataset_size is None + assert config.max_dataset_size == "all" def test_general_dataset_default_remains_uncapped() -> None: config = DatasetConfiguration(seeds=make_objectives(*(str(index) for index in range(8)))) - assert config.max_dataset_size is None - assert config.get_size_budget() == scenario_dataset_size_from_limit(None) + assert config.max_dataset_size == "all" + assert config.get_size_budget() == scenario_dataset_size_from_limit("all") assert config._apply_max_dataset_size(list(range(8))) == list(range(8)) @@ -148,31 +151,31 @@ def test_init_with_seeds_only(self) -> None: config = DatasetConfiguration(seeds=seeds) assert config._seeds == seeds assert config._seed_groups is None - assert config._dataset_names is None + assert config.dataset_names == [] def test_init_with_seed_groups_only(self, sample_seed_groups: list[SeedGroup]) -> None: config = DatasetConfiguration(seed_groups=sample_seed_groups) assert config._seed_groups == sample_seed_groups assert config._seeds is None - assert config._dataset_names is None - assert config.max_dataset_size is None + assert config.dataset_names == [] + assert config.max_dataset_size == "all" def test_init_with_dataset_names_only(self) -> None: config = DatasetConfiguration(dataset_names=["dataset1", "dataset2"]) - assert config._dataset_names == ["dataset1", "dataset2"] + assert config.dataset_names == ["dataset1", "dataset2"] assert config._seeds is None assert config._seed_groups is None def test_init_defaults_to_auto_fetch(self) -> None: config = DatasetConfiguration(dataset_names=["d1"]) - assert config._auto_fetch is True + assert config.fetch is DatasetFetchPolicy.IF_MISSING def test_init_auto_fetch_can_be_disabled(self) -> None: config = DatasetConfiguration(dataset_names=["d1"], auto_fetch=False) - assert config._auto_fetch is False + assert config.fetch is DatasetFetchPolicy.NEVER def test_init_with_two_sources_raises(self, sample_seed_groups: list[SeedGroup]) -> None: - with pytest.raises(ValueError, match="Only one of 'seeds', 'seed_groups', or 'dataset_names'"): + with pytest.raises(ValueError, match="Only one of"): DatasetConfiguration(seed_groups=sample_seed_groups, dataset_names=["d1"]) def test_init_with_three_sources_raises(self, sample_seed_groups: list[SeedGroup]) -> None: @@ -205,7 +208,7 @@ def test_init_copies_dataset_names_to_prevent_mutation(self) -> None: names = ["d1", "d2"] config = DatasetConfiguration(dataset_names=names) names.append("d3") - assert config._dataset_names == ["d1", "d2"] + assert config.dataset_names == ["d1", "d2"] def test_init_copies_seeds_to_prevent_mutation(self) -> None: seeds = make_objectives("a", "b") @@ -245,13 +248,12 @@ async def test_empty_inline_raises(self) -> None: async def test_raises_loudly_when_still_empty_after_fetch(self) -> None: config = DatasetAttackConfiguration(dataset_names=["d1"]) - with patch.object(config, "_fetch_dataset_async", new=AsyncMock()): - with pytest.raises(DatasetConstraintError, match="could not be loaded"): - await config.get_attack_seed_groups_async() + with pytest.raises(DatasetConstraintError, match="could not be loaded"): + await config.get_attack_seed_groups_async() async def test_raises_when_empty_and_auto_fetch_disabled(self) -> None: config = DatasetAttackConfiguration(dataset_names=["d1"], auto_fetch=False) - with pytest.raises(DatasetConstraintError, match="auto_fetch is disabled"): + with pytest.raises(DatasetConstraintError, match="prepare_async"): await config.get_attack_seed_groups_async() async def test_dataset_constraint_error_is_value_error(self) -> None: @@ -292,13 +294,18 @@ async def test_empty_raises(self) -> None: with pytest.raises(DatasetConstraintError): await config.get_attack_seed_groups_async() - async def test_auto_fetch_when_memory_empty(self, mock_memory: MagicMock) -> None: - mock_memory.get_seeds_async = AsyncMock(side_effect=[[], make_objectives("a")]) + async def test_explicit_prepare_then_read(self, mock_memory: MagicMock) -> None: + mock_memory.get_seeds_async = AsyncMock(return_value=make_objectives("a")) config = DatasetAttackConfiguration(dataset_names=["d1"]) - with patch.object(config, "_fetch_dataset_async", new=AsyncMock()) as mock_fetch: + dataset = SeedDataset(dataset_name="d1", seeds=[SeedObjective(value="a", dataset_name="d1")]) + fetcher = MagicMock() + fetcher.fetch_dataset_async = AsyncMock(return_value=dataset) + with patch(PROVIDER_PATCH_TARGET) as provider: + provider.get_providers_by_name_async = AsyncMock(return_value={"d1": fetcher}) + await config.prepare_async() groups = await config.get_attack_seed_groups_async() assert len(groups) == 1 - mock_fetch.assert_awaited_once_with(dataset_name="d1") + fetcher.fetch_dataset_async.assert_awaited_once() class TestGetAttackGroupsByDatasetAsync: @@ -355,58 +362,57 @@ def _build_attack_groups(self, seeds): assert await config.get_attack_seed_groups_async() == sentinel -class TestFetchDatasetAsync: - """``_fetch_dataset_async`` provider interaction.""" +class TestPrepareAsync: + """Preparation is the only provider and persistence boundary.""" async def test_unregistered_name_does_not_fetch(self, mock_memory: MagicMock) -> None: config = DatasetConfiguration(dataset_names=["d1"]) with patch(PROVIDER_PATCH_TARGET) as provider: - provider.get_all_dataset_names_async = AsyncMock(return_value=["other"]) - provider.fetch_datasets_async = AsyncMock() - await config._fetch_dataset_async(dataset_name="d1") - provider.fetch_datasets_async.assert_not_called() + provider.get_providers_by_name_async = AsyncMock(return_value={}) + with pytest.raises(DatasetConstraintError, match="Import them first"): + await config.prepare_async() mock_memory.add_seed_datasets_to_memory_async.assert_not_called() async def test_registered_name_fetches_and_adds(self, mock_memory: MagicMock) -> None: config = DatasetConfiguration(dataset_names=["d1"]) - datasets = [MagicMock()] + dataset = SeedDataset(dataset_name="d1", seeds=[SeedObjective(value="a", dataset_name="d1")]) + fetcher = MagicMock() + fetcher.fetch_dataset_async = AsyncMock(return_value=dataset) with patch(PROVIDER_PATCH_TARGET) as provider: - provider.get_all_dataset_names_async = AsyncMock(return_value=["d1"]) - provider.fetch_datasets_async = AsyncMock(return_value=datasets) - await config._fetch_dataset_async(dataset_name="d1") - provider.fetch_datasets_async.assert_awaited_once_with(dataset_names=["d1"]) + provider.get_providers_by_name_async = AsyncMock(return_value={"d1": fetcher}) + await config.prepare_async() + fetcher.fetch_dataset_async.assert_awaited_once() mock_memory.add_seed_datasets_to_memory_async.assert_awaited_once() async def test_enumeration_error_propagates(self, mock_memory: MagicMock) -> None: config = DatasetConfiguration(dataset_names=["d1"]) with patch(PROVIDER_PATCH_TARGET) as provider: - provider.get_all_dataset_names_async = AsyncMock(side_effect=RuntimeError("boom")) + provider.get_providers_by_name_async = AsyncMock(side_effect=RuntimeError("boom")) with pytest.raises(RuntimeError, match="boom"): - await config._fetch_dataset_async(dataset_name="d1") + await config.prepare_async() mock_memory.add_seed_datasets_to_memory_async.assert_not_called() - async def test_fetch_failure_chains_root_cause(self, mock_memory: MagicMock) -> None: + async def test_fetch_failure_propagates(self, mock_memory: MagicMock) -> None: config = DatasetAttackConfiguration(dataset_names=["d1"]) + fetcher = MagicMock() + fetcher.fetch_dataset_async = AsyncMock(side_effect=RuntimeError("boom")) with patch(PROVIDER_PATCH_TARGET) as provider: - provider.get_all_dataset_names_async = AsyncMock(side_effect=RuntimeError("boom")) - with pytest.raises(DatasetConstraintError, match="auto-fetch") as exc_info: - await config.get_attack_seed_groups_async() - assert isinstance(exc_info.value.__cause__, RuntimeError) + provider.get_providers_by_name_async = AsyncMock(return_value={"d1": fetcher}) + with pytest.raises(RuntimeError, match="boom"): + await config.prepare_async() + mock_memory.add_seed_datasets_to_memory_async.assert_not_awaited() async def test_read_only_resolution_does_not_fetch_or_persist(self, mock_memory: MagicMock) -> None: """Estimate resolution reports missing data without mutating central memory.""" config = DatasetAttackConfiguration(dataset_names=["d1"]) with ( patch(PROVIDER_PATCH_TARGET) as provider, - read_only_dataset_resolution(), - pytest.raises(DatasetConstraintError, match="read-only resolution"), + pytest.raises(DatasetConstraintError, match="prepare_async"), ): - provider.get_all_dataset_names_async = AsyncMock(return_value=["d1"]) - provider.fetch_datasets_async = AsyncMock() + provider.get_providers_by_name_async = AsyncMock() await config.get_attack_seed_groups_async() - provider.get_all_dataset_names_async.assert_not_awaited() - provider.fetch_datasets_async.assert_not_awaited() + provider.get_providers_by_name_async.assert_not_awaited() mock_memory.add_seed_datasets_to_memory_async.assert_not_awaited() @@ -570,7 +576,7 @@ def test_per_dataset_builds_one_child_per_name(self) -> None: config = CompoundDatasetAttackConfiguration.per_dataset(dataset_names=["d1", "d2"], max_dataset_size=4) assert len(config._configurations) == 2 assert [child.dataset_names for child in config._configurations] == [["d1"], ["d2"]] - assert all(child.max_dataset_size == 4 for child in config._configurations) + assert all(child.max_per_dataset == 4 for child in config._configurations) def test_size_caps_report_child_and_combined_limits(self) -> None: """Planning metadata explains independent child caps and the final compound cap.""" diff --git a/tests/unit/scenario/core/test_dataset_sources.py b/tests/unit/scenario/core/test_dataset_sources.py new file mode 100644 index 0000000000..f8b7bf4975 --- /dev/null +++ b/tests/unit/scenario/core/test_dataset_sources.py @@ -0,0 +1,430 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Named selection, explicit preparation, and durable storage contracts.""" + +from typing import Literal +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from sqlalchemy import event + +from pyrit.datasets import SeedDatasetProvider +from pyrit.memory import CentralMemory, MemoryInterface, SQLiteMemory +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import AllAvailableDatasetSize, BoundedDatasetSize, SeedDataset, SeedObjective, SeedOrigin +from pyrit.models.dataset_limit import DatasetLimit +from pyrit.scenario import ( + CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, + DatasetConfiguration, + DatasetFetchPolicy, + DatasetSource, +) +from pyrit.scenario.core.dataset_configuration import DatasetConstraintError, require_min_size + + +def dataset(name: str, count: int = 12) -> SeedDataset: + """Build an in-memory provider result.""" + return SeedDataset( + dataset_name=name, + seeds=[SeedObjective(value=f"{name}-{index}", dataset_name=name) for index in range(count)], + ) + + +@pytest.fixture +def memory() -> MagicMock: + """Memory with three complete fixture populations.""" + populations = {name: dataset(name).seeds for name in ("a", "b", "c")} + instance = MagicMock(spec=MemoryInterface) + instance.get_seed_dataset_names_async = AsyncMock(return_value=list(populations)) + instance.get_seeds_async = AsyncMock( + side_effect=lambda *, dataset_name, **kwargs: populations.get(dataset_name, []) + ) + instance.add_seed_datasets_to_memory_async = AsyncMock() + return instance + + +@pytest.mark.parametrize("total, expected", [(None, 15), (5, 5), (30, 15)]) +async def test_source_and_total_limits(*, memory: MagicMock, total: int | None, expected: int) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name=name) for name in ("a", "b", "c")], max_total=total) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + groups = await config.get_attack_groups_by_dataset_async() + assert sum(map(len, groups.values())) == expected + assert all(len(population) <= 5 for population in groups.values()) + assert len(await config.get_attack_seed_groups_async()) == expected + assert len(await config.get_attack_seed_groups_async(apply_sampling=False)) == 36 + assert config.get_size_budget() == BoundedDatasetSize(value=expected) + + +async def test_explicit_source_all_is_not_inherited(memory: MagicMock) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="a"), DatasetSource(name="b", max_size="all")]) + assert config.get_size_budget() == AllAvailableDatasetSize() + assert config.with_overrides(max_total=9).get_size_budget() == BoundedDatasetSize(value=9) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + groups = await config.get_attack_groups_by_dataset_async() + assert {name: len(items) for name, items in groups.items()} == {"a": 5, "b": 12} + + +@pytest.mark.parametrize("limit", [0, -1, True, 1.5]) +def test_source_rejects_invalid_limits(limit: int) -> None: + with pytest.raises(DatasetConstraintError, match="positive integer"): + DatasetSource(name="a", max_size=limit) + with pytest.raises(DatasetConstraintError, match="positive integer"): + DatasetAttackConfiguration(max_total=limit) + + +def test_source_names_and_alias_conflicts() -> None: + with pytest.raises(DatasetConstraintError, match="non-empty"): + DatasetSource(name=" ") + with pytest.raises(DatasetConstraintError, match="Duplicate"): + DatasetAttackConfiguration(sources=[DatasetSource(name="a"), DatasetSource(name="a")]) + with pytest.raises(ValueError, match="Only one"): + DatasetAttackConfiguration(sources=[DatasetSource(name="a")], dataset_names=["a"]) + with pytest.raises(ValueError, match="only one"): + DatasetAttackConfiguration(max_total=3, max_dataset_size=4) + + +@pytest.mark.parametrize("limit", [None, "", "default", "all", 3]) +def test_duplicate_total_arguments_raise_even_when_equal(limit: DatasetLimit) -> None: + with pytest.raises(ValueError, match="only one.*max_dataset_size.*max_total"): + DatasetAttackConfiguration(max_total=limit, max_dataset_size=limit) + + +@pytest.mark.parametrize("policy", list(DatasetFetchPolicy)) +def test_duplicate_fetch_arguments_raise_even_when_equal(policy: DatasetFetchPolicy) -> None: + with pytest.raises(ValueError, match="only one.*auto_fetch.*fetch"): + DatasetAttackConfiguration(fetch=policy, auto_fetch=policy is DatasetFetchPolicy.IF_MISSING) + + +@pytest.mark.parametrize("compound", [False, True]) +async def test_preparation_checks_names_without_reading_seeds_async(*, memory: MagicMock, compound: bool) -> None: + sources = [DatasetSource(name="a"), DatasetSource(name="b", fetch=DatasetFetchPolicy.NEVER)] + config = ( + CompoundDatasetAttackConfiguration( + configurations=[DatasetAttackConfiguration(sources=[source]) for source in sources] + ) + if compound + else DatasetAttackConfiguration(sources=sources) + ) + memory.get_seeds_async.side_effect = AssertionError("Preparation must not load seed contents") + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "get_providers_by_name_async") as lookup, + ): + await config.prepare_async() + memory.get_seed_dataset_names_async.assert_awaited_once() + memory.get_seeds_async.assert_not_awaited() + lookup.assert_not_awaited() + memory.add_seed_datasets_to_memory_async.assert_not_awaited() + + +@pytest.mark.parametrize("total", [None, "", "default", "all", 3, 9]) +@pytest.mark.parametrize("source_limit, population", [("all", 12), (None, 5), ("default", 5), (5, 5)]) +async def test_legacy_total_has_canonical_defaults_async( + *, memory: MagicMock, total: DatasetLimit, source_limit: DatasetLimit, population: int +) -> None: + with pytest.warns(DeprecationWarning): + legacy = DatasetAttackConfiguration( + sources=[DatasetSource(name="a")], max_dataset_size=total, max_per_dataset=source_limit + ) + canonical = DatasetAttackConfiguration( + sources=[DatasetSource(name="a")], max_total=total, max_per_dataset=source_limit + ) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + for config in (legacy, canonical): + assert len(await config.get_attack_seed_groups_async()) == min( + population, total if isinstance(total, int) else population + ) + assert legacy.get_size_budget() == canonical.get_size_budget() + + +def test_deprecated_total_preserves_default_source_cap() -> None: + with pytest.warns(DeprecationWarning): + legacy = DatasetAttackConfiguration(sources=[DatasetSource(name="a")], max_dataset_size=10) + canonical = DatasetAttackConfiguration(sources=[DatasetSource(name="a")], max_total=10) + assert legacy.max_per_dataset == canonical.max_per_dataset == 5 + assert legacy.get_size_budget() == canonical.get_size_budget() == BoundedDatasetSize(value=5) + + +@pytest.mark.parametrize(("total", "expected"), [(None, 5), ("", 5), ("default", 5), ("all", 8), (3, 3), (9, 8)]) +async def test_inline_total_alias_parity_async(*, total: DatasetLimit, expected: int) -> None: + seeds = [SeedObjective(value=str(index)) for index in range(8)] + with pytest.warns(DeprecationWarning): + legacy = DatasetAttackConfiguration(seeds=seeds, max_dataset_size=total) + canonical = DatasetAttackConfiguration(seeds=seeds, max_total=total) + assert legacy.max_per_dataset == canonical.max_per_dataset == "all" + for config in (legacy, canonical): + assert len(await config.get_attack_seed_groups_async()) == expected + + +@pytest.mark.parametrize("configuration_class", [DatasetConfiguration, DatasetAttackConfiguration]) +def test_inline_source_caps_raise_on_construction_and_override( + configuration_class: type[DatasetConfiguration], +) -> None: + seeds = [SeedObjective(value="inline")] + with pytest.raises(DatasetConstraintError, match="Inline.*max_total"): + configuration_class(seeds=seeds, max_per_dataset=2) + config = configuration_class(seeds=seeds, max_per_dataset=None) + with pytest.raises(DatasetConstraintError, match="Inline.*max_total"): + config.with_overrides(max_per_dataset=2) + assert config.with_overrides(max_total=None).max_total == config.max_total + assert config.with_overrides(max_total="all").max_total == "all" + + +def test_total_only_source_caps_raise_but_explicit_all_is_unlimited() -> None: + config = DatasetAttackConfiguration( + sources=[DatasetSource(name="ingredient", max_size="all")], sampling_scope="total_only" + ) + with pytest.raises(DatasetConstraintError, match="Total-only.*max_total"): + DatasetAttackConfiguration(sources=[DatasetSource(name="ingredient", max_size=2)], sampling_scope="total_only") + with pytest.raises(DatasetConstraintError, match="Total-only.*max_total"): + config.with_overrides(max_per_dataset=2) + assert config.source_limit("ingredient") == "all" + + +@pytest.mark.parametrize(("scope", "expected"), [("per_dataset", 10), ("total_only", 17)]) +async def test_sampling_scope_controls_default_caps_async( + *, memory: MagicMock, scope: Literal["per_dataset", "total_only"], expected: int +) -> None: + config = DatasetAttackConfiguration( + sources=[DatasetSource(name="a"), DatasetSource(name="b")], sampling_scope=scope, max_total=17 + ) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + assert len(await config.get_attack_seed_groups_async()) == expected + assert config.get_size_budget() == BoundedDatasetSize(value=expected) + changed = config.with_overrides(max_total=3) + assert changed.sampling_scope == scope + assert config.max_total == 17 + assert changed.get_size_budget() == BoundedDatasetSize(value=3) + + +def test_invalid_sampling_scope_raises() -> None: + with pytest.raises(DatasetConstraintError, match="sampling_scope"): + DatasetAttackConfiguration(sampling_scope="invalid") # type: ignore[arg-type] + + +@pytest.mark.parametrize("default", [None, "", "default"]) +async def test_default_limits_match_omission_async(*, memory: MagicMock, default: DatasetLimit) -> None: + config = DatasetAttackConfiguration( + sources=[DatasetSource(name="a", max_size=default), DatasetSource(name="b")], + max_per_dataset=default, + max_total=default, + ) + assert config.max_per_dataset == 5 + assert config.max_total == "all" + assert config.sources[0].max_size == "default" + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + groups = await config.get_attack_groups_by_dataset_async() + assert {name: len(population) for name, population in groups.items()} == {"a": 5, "b": 5} + inline = DatasetAttackConfiguration(seeds=dataset("inline").seeds, max_total=default) + assert len(await inline.get_attack_seed_groups_async()) == 5 + + +@pytest.mark.parametrize("default", [None, "", "default"]) +async def test_default_overrides_preserve_existing_caps_async(*, memory: MagicMock, default: DatasetLimit) -> None: + config = DatasetAttackConfiguration( + sources=[DatasetSource(name="a", max_size=default)], max_per_dataset=4, max_total=3 + ) + copied = config.with_overrides(max_per_dataset=default, max_total=default) + assert copied.max_per_dataset == 4 + assert copied.max_total == 3 + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + assert len(await copied.get_attack_seed_groups_async()) == 3 + assert len(await copied.with_overrides(max_total="all").get_attack_seed_groups_async()) == 4 + unlimited = copied.with_overrides(max_per_dataset="all", max_total="all") + assert len(await unlimited.get_attack_seed_groups_async()) == 12 + assert unlimited.max_per_dataset == unlimited.max_total == "all" + + +@pytest.mark.parametrize("default", [None, "", "default"]) +def test_limit_setters_use_constructor_defaults(default: DatasetLimit) -> None: + config = DatasetAttackConfiguration(seeds=dataset("inline").seeds, max_total=2) + config.max_total = default + assert config.max_total == 5 + with pytest.warns(DeprecationWarning): + config.max_dataset_size = "all" + assert config.max_total == "all" + with pytest.warns(DeprecationWarning): + config.max_dataset_size = default + assert config.max_total == 5 + named = DatasetAttackConfiguration(sources=[DatasetSource(name="a")], max_per_dataset="all") + named.max_per_dataset = default + assert named.max_per_dataset == 5 + + +@pytest.mark.parametrize("default", [None, "", "default"]) +async def test_compound_default_and_all_preserve_child_caps_async(default: DatasetLimit) -> None: + child = DatasetAttackConfiguration(seeds=dataset("inline").seeds, max_total=3) + compound = CompoundDatasetAttackConfiguration(configurations=[child], max_total=default) + assert compound.max_total == "all" + assert len(await compound.get_attack_seed_groups_async()) == 3 + capped = compound.with_overrides(max_total=2) + assert len(await capped.with_overrides(max_total=default).get_attack_seed_groups_async()) == 2 + assert len(await capped.with_overrides(max_total="all").get_attack_seed_groups_async()) == 3 + + +def test_compound_source_override_rejects_inline_child_cap() -> None: + config = CompoundDatasetAttackConfiguration( + configurations=[DatasetAttackConfiguration(seeds=[SeedObjective(value="inline")])] + ) + with pytest.raises(DatasetConstraintError, match="Inline.*max_total"): + config.with_overrides(max_per_dataset=2) + + +def test_overrides_keep_class_and_do_not_mutate_containers() -> None: + class CustomConfiguration(DatasetAttackConfiguration): + pass + + original = CustomConfiguration(sources=[DatasetSource(name="a")], filters={"harm_categories": ["first"]}) + changed = original.with_overrides(max_total=2, filters={"harm_categories": ["second"]}) + assert type(changed) is CustomConfiguration + assert changed.sources == original.sources + assert original.max_total == "all" + assert original.filters == {"harm_categories": ["first"]} + compound = CompoundDatasetAttackConfiguration(configurations=[original]) + copied = compound.with_overrides(fetch=DatasetFetchPolicy.NEVER, filters={"harm_categories": ["second"]}) + assert copied._configurations[0] is not original + assert original.fetch is DatasetFetchPolicy.IF_MISSING + assert copied._configurations[0].fetch is DatasetFetchPolicy.NEVER + assert original.filters == {"harm_categories": ["first"]} + + +async def test_compound_validates_full_population_before_sampling(memory: MagicMock) -> None: + config = CompoundDatasetAttackConfiguration( + configurations=[DatasetAttackConfiguration(sources=[DatasetSource(name="a")], max_total=1)], + validators=[require_min_size(12)], + ) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + assert len(await config.get_attack_seed_groups_async()) == 1 + + +@pytest.mark.parametrize("compound", [False, True]) +@pytest.mark.parametrize("policy", [DatasetFetchPolicy.NEVER, DatasetFetchPolicy.IF_MISSING]) +async def test_preflight_checks_last_missing_source_before_any_fetch( + *, memory: MagicMock, compound: bool, policy: DatasetFetchPolicy +) -> None: + first = DatasetSource(name="registered") + last = DatasetSource(name="memory_only", fetch=policy) + config = ( + CompoundDatasetAttackConfiguration( + configurations=[DatasetAttackConfiguration(sources=[source]) for source in (first, last)] + ) + if compound + else DatasetAttackConfiguration(sources=[first, last]) + ) + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async = AsyncMock(return_value=dataset("registered")) + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"registered": provider}), + pytest.raises(DatasetConstraintError, match="Import"), + ): + await config.prepare_async() + provider.fetch_dataset_async.assert_not_awaited() + memory.add_seed_datasets_to_memory_async.assert_not_awaited() + + +async def test_filter_miss_never_fetches(memory: MagicMock) -> None: + memory.get_seeds_async.side_effect = lambda **kwargs: [] if "harm_categories" in kwargs else dataset("a").seeds + config = DatasetAttackConfiguration(sources=[DatasetSource(name="a")], filters={"harm_categories": ["absent"]}) + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "get_providers_by_name_async") as lookup, + ): + await config.prepare_async() + with pytest.raises(DatasetConstraintError, match="none match"): + await config.get_attack_seed_groups_async() + lookup.assert_not_awaited() + memory.add_seed_datasets_to_memory_async.assert_not_awaited() + + +async def test_provider_failure_is_not_cached_or_persisted(memory: MagicMock) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="fresh")]) + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async.side_effect = RuntimeError("provider failed") + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"fresh": provider}), + ): + for _ in range(2): + with pytest.raises(RuntimeError, match="provider failed"): + await config.prepare_async() + assert provider.fetch_dataset_async.await_count == 2 + memory.add_seed_datasets_to_memory_async.assert_not_awaited() + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("origin", list(SeedOrigin)) +async def test_memory_only_reuse_and_new_run_rechecks_presence( + *, sqlite_instance: SQLiteMemory, origin: SeedOrigin +) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="stored")]) + seeds = dataset("stored").seeds + for seed in seeds: + seed.origin = origin + await sqlite_instance.add_seeds_to_memory_async(seeds=seeds, added_by="test") + with patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={}) as lookup: + await config.prepare_async() + assert len(await config.get_attack_seed_groups_async()) == 5 + lookup.assert_not_awaited() + await sqlite_instance.remove_seeds_from_memory_async(dataset_name="stored") + with pytest.raises(DatasetConstraintError, match="no provider"): + await config.prepare_async() + lookup.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_preparation_persists_then_reuses(sqlite_instance: SQLiteMemory) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="fresh", max_size=2)]) + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async = AsyncMock(return_value=dataset("fresh")) + with patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"fresh": provider}): + await config.prepare_async() + assert len(await config.get_attack_seed_groups_async()) == 2 + assert len(await sqlite_instance.get_seeds_async(dataset_name="fresh")) == 12 + await config.prepare_async() + provider.fetch_dataset_async.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_insert_failure_rolls_back_complete_dataset(sqlite_instance: SQLiteMemory) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="fresh")]) + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async = AsyncMock(return_value=dataset("fresh", count=2)) + inserted = 0 + + def fail_second_insert(mapper: object, connection: object, target: SeedEntry) -> None: + nonlocal inserted + inserted += 1 + if inserted == 2: + raise RuntimeError("injected insert failure") + + event.listen(SeedEntry, "before_insert", fail_second_insert) + try: + with ( + patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"fresh": provider}), + pytest.raises(RuntimeError, match="injected insert failure"), + ): + await config.prepare_async() + finally: + event.remove(SeedEntry, "before_insert", fail_second_insert) + assert await sqlite_instance.get_seeds_async(dataset_name="fresh") == [] + + +@pytest.mark.parametrize("name, seeds", [("wrong", [SeedObjective(value="x", dataset_name="wrong")]), ("fresh", [])]) +async def test_invalid_provider_result_never_persists( + *, memory: MagicMock, name: str, seeds: list[SeedObjective] +) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="fresh")]) + provider = MagicMock(spec=SeedDatasetProvider) + result = dataset(name) + result.seeds = seeds + provider.fetch_dataset_async = AsyncMock(return_value=result) + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"fresh": provider}), + pytest.raises(DatasetConstraintError, match="non-empty seeds"), + ): + await config.prepare_async() + memory.add_seed_datasets_to_memory_async.assert_not_awaited() diff --git a/tests/unit/scenario/core/test_scenario.py b/tests/unit/scenario/core/test_scenario.py index bfef6b0280..01d86d6287 100644 --- a/tests/unit/scenario/core/test_scenario.py +++ b/tests/unit/scenario/core/test_scenario.py @@ -10,6 +10,7 @@ import pytest +from pyrit.datasets import SeedDatasetProvider from pyrit.executor.attack import PromptSendingAttack, RedTeamingAttack from pyrit.executor.attack.core import AttackExecutorResult from pyrit.memory import CentralMemory, MemoryInterface @@ -20,6 +21,7 @@ AttackSeedGroup, ComponentIdentifier, ScenarioRunState, + SeedDataset, SeedObjective, SeedPrompt, ) @@ -28,6 +30,7 @@ from pyrit.scenario import ( DatasetAttackConfiguration, DatasetConfiguration, + DatasetSource, ScenarioIdentifier, ScenarioResult, ) @@ -983,6 +986,9 @@ async def test_initialize_async_with_empty_techniques_and_baseline(self, mock_ob # Create a mock dataset config with seed groups mock_dataset_config = MagicMock(spec=DatasetAttackConfiguration) + mock_dataset_config.with_overrides.return_value = mock_dataset_config + mock_dataset_config.sources = () + mock_dataset_config.max_total = "all" mock_dataset_config.get_attack_groups_by_dataset_async.return_value = { "default": [ AttackSeedGroup(seeds=[SeedObjective(value="test objective 1")]), @@ -1016,6 +1022,9 @@ async def test_baseline_only_execution_runs_successfully(self, mock_objective_ta # Create a mock dataset config with seed groups mock_dataset_config = MagicMock(spec=DatasetAttackConfiguration) + mock_dataset_config.with_overrides.return_value = mock_dataset_config + mock_dataset_config.sources = () + mock_dataset_config.max_total = "all" mock_dataset_config.get_attack_groups_by_dataset_async.return_value = { "default": [AttackSeedGroup(seeds=[SeedObjective(value="test objective 1")])] } @@ -1051,6 +1060,9 @@ async def test_empty_techniques_without_baseline_allows_initialization(self, moc ) mock_dataset_config = MagicMock(spec=DatasetConfiguration) + mock_dataset_config.with_overrides.return_value = mock_dataset_config + mock_dataset_config.sources = () + mock_dataset_config.max_total = "all" # None techniques with no baseline: _get_atomic_attacks_async returns [] scenario.set_params_from_args( @@ -1090,6 +1102,9 @@ async def test_standalone_baseline_uses_dataset_config_seeds(self, mock_objectiv ] mock_dataset_config = MagicMock(spec=DatasetAttackConfiguration) + mock_dataset_config.with_overrides.return_value = mock_dataset_config + mock_dataset_config.sources = () + mock_dataset_config.max_total = "all" mock_dataset_config.get_attack_groups_by_dataset_async.return_value = {"default": expected_seeds} scenario.set_params_from_args( @@ -1398,6 +1413,58 @@ def _make_config(self): seed_groups = [SeedGroup(seeds=[SeedObjective(value=f"obj{i}")]) for i in range(10)] return DatasetAttackConfiguration(seed_groups=seed_groups, max_dataset_size=3) + @pytest.mark.parametrize("resume_state", ["valid", "missing-id", "missing-dataset", "incompatible"]) + async def test_named_source_preparation_and_resume(self, mock_objective_target, resume_state: str) -> None: + config = DatasetAttackConfiguration(sources=[DatasetSource(name="fixture", max_size=3)]) + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async.return_value = SeedDataset( + dataset_name="fixture", + seeds=[SeedObjective(value=f"fixture-{index}", dataset_name="fixture") for index in range(10)], + ) + scenario = self._StrategyScenario(name="Named sources", version=1) + scenario.set_params_from_args(args={"objective_target": mock_objective_target, "dataset_config": config}) + with ( + patch.object(config, "prepare_async", wraps=config.prepare_async) as prepare, + patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={"fixture": provider}), + ): + await scenario.initialize_async() + prepare.assert_awaited_once() + provider.fetch_dataset_async.assert_awaited_once() + baseline, strategy = scenario._atomic_attacks + assert baseline.seed_groups == strategy.seed_groups + assert len(strategy.seed_groups) == 3 + original_plan = scenario._build_run_plan() + + if resume_state == "missing-dataset": + await scenario._memory.remove_seeds_from_memory_async(dataset_name="fixture") + result_id = ( + "00000000-0000-0000-0000-000000000000" if resume_state == "missing-id" else scenario._scenario_result_id + ) + resumed = self._StrategyScenario( + name="Named sources", + version=2 if resume_state == "incompatible" else 1, + scenario_result_id=result_id, + ) + resumed.set_params_from_args(args={"objective_target": mock_objective_target, "dataset_config": config}) + with ( + patch.object(config, "prepare_async", side_effect=AssertionError("Resume must not prepare")) as prepare, + patch.object( + SeedDatasetProvider, "get_providers_by_name_async", side_effect=AssertionError("No lookup") + ) as lookup, + patch( + "pyrit.scenario.core.dataset_configuration.random.sample", side_effect=AssertionError("No resampling") + ) as sample, + ): + if resume_state == "valid": + await resumed.initialize_async() + assert resumed._build_run_plan() == original_plan + else: + with pytest.raises(ValueError): + await resumed.initialize_async() + prepare.assert_not_awaited() + lookup.assert_not_awaited() + sample.assert_not_called() + async def test_resume_reconstructs_persisted_subset_without_resampling(self, mock_objective_target): config = self._make_config() @@ -1682,6 +1749,9 @@ async def test_resume_with_composite_scorer_async( aggregator=aggregator, scorers=[SubStringScorer(substring=value) for value in ("a", "b")] ) dataset_config = MagicMock(spec=DatasetAttackConfiguration) + dataset_config.with_overrides.return_value = dataset_config + dataset_config.sources = () + dataset_config.max_total = "all" dataset_config.get_attack_groups_by_dataset_async.return_value = { "default": [AttackSeedGroup(seeds=[SeedObjective(value="test objective")])] } diff --git a/tests/unit/scenario/foundry/test_red_team_agent.py b/tests/unit/scenario/foundry/test_red_team_agent.py index b8a4de4675..ad15602a47 100644 --- a/tests/unit/scenario/foundry/test_red_team_agent.py +++ b/tests/unit/scenario/foundry/test_red_team_agent.py @@ -130,6 +130,9 @@ def mock_memory_seed_groups(): def mock_dataset_config(mock_memory_seed_groups): """Create a mock dataset config that returns the seed groups.""" mock_config = MagicMock(spec=DatasetAttackConfiguration) + mock_config.with_overrides.return_value = mock_config + mock_config.sources = () + mock_config.max_total = "all" mock_config.get_attack_seed_groups_async = AsyncMock(return_value=mock_memory_seed_groups) mock_config.dataset_names = ["foundry_red_team"] return mock_config @@ -341,7 +344,7 @@ async def test_init_raises_exception_when_no_datasets_available(self, mock_objec # Error should occur during initialize_async when it resolves seed groups. # Neutralize the provider fetch so the empty-memory path raises loudly instead of fetching. with patch( - "pyrit.scenario.core.dataset_configuration.DatasetConfiguration._fetch_dataset_async", + "pyrit.scenario.core.dataset_configuration.DatasetConfiguration.prepare_async", new_callable=AsyncMock, ): scenario.set_params_from_args( diff --git a/tests/unit/scenario/garak/test_api_key.py b/tests/unit/scenario/garak/test_api_key.py index 87026c43f6..0a93b2389c 100644 --- a/tests/unit/scenario/garak/test_api_key.py +++ b/tests/unit/scenario/garak/test_api_key.py @@ -4,6 +4,7 @@ """Tests for the Garak API-key scenario.""" from pathlib import Path +from typing import Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -107,7 +108,7 @@ async def test_uncapped_configuration_renders_full_corpus( scenario=scenario, target=mock_objective_target, corpus_seeds=corpus_seeds, - dataset_config=ApiKeyDatasetConfiguration(dataset_names=ApiKey.required_datasets(), max_dataset_size=None), + dataset_config=ApiKeyDatasetConfiguration(dataset_names=ApiKey.required_datasets(), max_dataset_size="all"), ) assert {name: len(groups) for name, groups in _objectives(scenario).items()} == { @@ -234,9 +235,12 @@ async def test_unsupported_dataset_configuration_raises( scenario=ApiKey(), target=mock_objective_target, corpus_seeds=corpus_seeds, dataset_config=config ) - @pytest.mark.parametrize("size", [1, 7, 20, 348, None]) + @pytest.mark.parametrize("size", [1, 7, 20, 348, None, "all"]) async def test_launch_and_estimate_use_standard_dataset_size( - self, size: int | None, mock_objective_target: PromptTarget, corpus_seeds: dict[str, list[Seed]] + self, + size: int | Literal["all"] | None, + mock_objective_target: PromptTarget, + corpus_seeds: dict[str, list[Seed]], ) -> None: scenario = ApiKey() args = ScenarioConfigurationResolver.resolve_configuration( @@ -255,10 +259,15 @@ async def test_launch_and_estimate_use_standard_dataset_size( estimate = await scenario.get_run_size_estimate_async(target_is_configured=True) await scenario.initialize_async() - expected = size or 20 - assert estimate.status is ScenarioRunSizeEstimateStatus.Approximate - assert estimate.total_attack_count == expected - assert estimate.estimated_attack_count == expected + total = None if size == "all" else size or 20 + expected = 348 if total is None else total + assert estimate.status is ( + ScenarioRunSizeEstimateStatus.Approximate + if total is not None + else ScenarioRunSizeEstimateStatus.Unavailable + ) + assert estimate.total_attack_count == total + assert estimate.estimated_attack_count == total assert estimate.minimum_attack_count is None assert estimate.maximum_attack_count is None assert all( @@ -274,12 +283,10 @@ async def test_launch_and_estimate_use_standard_dataset_size( for dataset in estimate.datasets: assert dataset.kind == "synthesized" - assert len(dataset.configured_caps) == 1 - cap = dataset.configured_caps[0] - assert cap.label == "combined configuration cap" - assert cap.count == expected - assert cap.configured_on == "configuration" - assert cap.dataset_name == dataset.name + assert [(cap.label, cap.count, cap.configured_on) for cap in dataset.configured_caps] == ( + [("combined configuration cap", total, "configuration")] if total is not None else [] + ) + assert all(cap.dataset_name == dataset.name for cap in dataset.configured_caps) @pytest.mark.parametrize("size", [None, 3]) @pytest.mark.parametrize("technique", [ApiKeyTechnique.GetKey, ApiKeyTechnique.CompleteKey]) diff --git a/tests/unit/scenario/garak/test_divergence.py b/tests/unit/scenario/garak/test_divergence.py index dfc1f51956..bd07e847fa 100644 --- a/tests/unit/scenario/garak/test_divergence.py +++ b/tests/unit/scenario/garak/test_divergence.py @@ -5,6 +5,7 @@ from collections import Counter from pathlib import Path +from typing import Literal from unittest.mock import AsyncMock, patch import pytest @@ -124,7 +125,9 @@ async def test_selection_runs_each_seed_once( await _initialize_async( scenario=scenario, corpus=corpus, - config=DivergenceDatasetConfiguration(dataset_names=["garak_divergence"], max_dataset_size=None), + config=DivergenceDatasetConfiguration( + dataset_names=["garak_divergence"], max_per_dataset="all", max_total="all" + ), techniques=techniques, ) assert len(_groups(scenario)) == 36 @@ -144,8 +147,8 @@ async def test_selection_runs_each_seed_once( assert all(group.prompts[0].metadata["repeat_word"] == condition.text for group in attack.seed_groups) assert attack.attack_technique.attack._objective_scorer is scenario._objective_scorer - @pytest.mark.parametrize("size", [None, 1, 7, 36]) - async def test_runtime_budget_matches_estimate(self, corpus: list[Seed], size: int | None) -> None: + @pytest.mark.parametrize("size", [None, 1, 7, 36, "all"]) + async def test_runtime_budget_matches_estimate(self, corpus: list[Seed], size: int | Literal["all"] | None) -> None: scenario = Divergence() args = ScenarioConfigurationResolver.resolve_configuration( scenario_name="garak.divergence", @@ -162,8 +165,9 @@ async def test_runtime_budget_matches_estimate(self, corpus: list[Seed], size: i ): estimate = await scenario.get_run_size_estimate_async(target_is_configured=True) await scenario.initialize_async() - assert estimate.estimated_attack_count == (size or 10) - assert len(_groups(scenario)) == (size or 10) + total = None if size == "all" else size or 10 + assert estimate.estimated_attack_count == total + assert len(_groups(scenario)) == (total if total is not None else 36) async def test_resume_preserves_sample_and_group_identity(self, corpus: list[Seed]) -> None: scenario = Divergence() @@ -266,7 +270,9 @@ async def test_full_run_scores_and_persists_each_expectation( scenario=scenario, corpus=corpus, target=target, - config=DivergenceDatasetConfiguration(dataset_names=["garak_divergence"], max_dataset_size=None), + config=DivergenceDatasetConfiguration( + dataset_names=["garak_divergence"], max_per_dataset="all", max_total="all" + ), ) word_by_prompt = {seed.value: seed.metadata["repeat_word"] for seed in corpus} diff --git a/tests/unit/scenario/garak/test_encoding.py b/tests/unit/scenario/garak/test_encoding.py index e8d8e015e0..f6ef5ea8f1 100644 --- a/tests/unit/scenario/garak/test_encoding.py +++ b/tests/unit/scenario/garak/test_encoding.py @@ -62,6 +62,9 @@ def mock_attack_seed_groups(mock_memory_seeds): def mock_dataset_config(mock_attack_seed_groups): """Create a mock dataset config that returns the seed attack groups.""" mock_config = MagicMock(spec=EncodingDatasetConfiguration) + mock_config.with_overrides.return_value = mock_config + mock_config.sources = () + mock_config.max_total = "all" mock_config.get_attack_seed_groups_async = AsyncMock(return_value=mock_attack_seed_groups) mock_config.dataset_names = ["garak_slur_terms_en", "garak_web_html_js"] return mock_config @@ -137,7 +140,7 @@ async def test_init_raises_exception_when_no_datasets_available(self, mock_objec # Disable the provider fallback so memory stays empty and the scenario raises. scenario = Encoding(objective_scorer=mock_objective_scorer) - with patch.object(EncodingDatasetConfiguration, "_fetch_dataset_async", new_callable=AsyncMock): + with patch.object(DatasetConfiguration, "prepare_async", new_callable=AsyncMock): # Error should occur during initialize_async when _get_atomic_attacks_async resolves seed prompts scenario.set_params_from_args(args={"objective_target": mock_objective_target}) with pytest.raises(DatasetConstraintError, match="could not be loaded"): @@ -675,7 +678,7 @@ def test_encoding_dataset_config_can_be_initialized_with_dataset_names(self): max_dataset_size=5, ) - assert config._dataset_names == ["garak_slur_terms_en", "garak_web_html_js"] + assert config.dataset_names == ["garak_slur_terms_en", "garak_web_html_js"] assert config.max_dataset_size == 5 diff --git a/tests/unit/scenario/garak/test_figstep.py b/tests/unit/scenario/garak/test_figstep.py index b994d3e14f..06d5988256 100644 --- a/tests/unit/scenario/garak/test_figstep.py +++ b/tests/unit/scenario/garak/test_figstep.py @@ -171,6 +171,7 @@ async def test_named_pro_dataset_selects_pro_attack(self, mock_objective_target, ) with ( patch("pyrit.prompt_target.common.target_requirements.TargetRequirements.validate"), + patch.object(config, "prepare_async", new_callable=AsyncMock), patch.object( DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", diff --git a/tests/unit/scenario/garak/test_latent_injection.py b/tests/unit/scenario/garak/test_latent_injection.py index 940b84d85a..8af5a269b2 100644 --- a/tests/unit/scenario/garak/test_latent_injection.py +++ b/tests/unit/scenario/garak/test_latent_injection.py @@ -84,18 +84,27 @@ def _ids(scenario: LatentInjection) -> dict[str, list[str]]: class TestLatentDefaults: @pytest.mark.parametrize( ("kwargs", "expected_cap"), - [({}, 92), ({"max_dataset_size": None}, None), ({"max_dataset_size": 23}, 23)], + [ + ({}, 92), + ({"max_dataset_size": None}, 92), + ({"max_total": "default"}, 92), + ({"max_total": ""}, 92), + ({"max_dataset_size": "all"}, "all"), + ({"max_total": "all"}, "all"), + ({"max_dataset_size": 23}, 23), + ], ) async def test_configuration_defaults_resolve_with_family_coverage_async( - self, *, kwargs: dict[str, int | None], expected_cap: int | None + self, *, kwargs: dict[str, int | str | None], expected_cap: int | str ) -> None: config = _config(**kwargs) + await config.prepare_async() groups = await config.get_attack_seed_groups_async() assert config.max_dataset_size == expected_cap assert len(config.coverage_keys) == 23 assert {config._coverage_key(group) for group in groups} == set(config.coverage_keys) full = await config.get_attack_seed_groups_async(apply_sampling=False) - assert len(groups) == (min(expected_cap, len(full)) if expected_cap is not None else len(full)) + assert len(groups) == (min(expected_cap, len(full)) if isinstance(expected_cap, int) else len(full)) assert len(full) > 92 async def test_default_population_budget_and_estimate_async(self) -> None: @@ -112,9 +121,10 @@ async def test_default_population_budget_and_estimate_async(self) -> None: assert {parameter.name for parameter in scenario.additional_parameters()} == {"families"} async def test_all_families_and_separators_async(self) -> None: - config = _config(families=LatentInjectionDatasetConfiguration.FAMILIES, max_dataset_size=None) + config = _config(families=LatentInjectionDatasetConfiguration.FAMILIES, max_dataset_size="all") scenario = LatentInjection(harm_scorer=SubStringScorer(substring="harm")) await _initialize_async(scenario, dataset_config=config, scenario_techniques=[LatentInjectionTechnique.ALL]) + config = scenario._dataset_config assert {key[0] for key in config.coverage_keys} == set(config.FAMILIES) assert len(scenario._atomic_attacks) == len(config.coverage_keys) * 14 for technique in LatentInjectionTechnique.expand([LatentInjectionTechnique.ALL]): @@ -131,8 +141,10 @@ async def test_all_families_and_separators_async(self) -> None: async def test_auto_fetch_false_and_wrong_configs_async(self) -> None: config = _config(auto_fetch=False) - with patch.object(config, "_fetch_dataset_async") as fetch: - with pytest.raises(DatasetConstraintError, match="auto_fetch is disabled"): + with patch( + "pyrit.datasets.seed_datasets.seed_dataset_provider.SeedDatasetProvider.get_providers_by_name_async" + ) as fetch: + with pytest.raises(DatasetConstraintError, match="fetch is 'never'"): await _initialize_async(LatentInjection(), dataset_config=config) fetch.assert_not_called() with pytest.raises(DatasetConstraintError, match="only supports"): @@ -150,7 +162,7 @@ def test_invalid_families(self, families: list[str]) -> None: @pytest.mark.parametrize("cap", [0, -1]) def test_invalid_constructor_budget(self, cap: int) -> None: - with pytest.raises(ValueError, match="max_dataset_size"): + with pytest.raises(ValueError, match="max_total"): _config(max_dataset_size=cap) @pytest.mark.parametrize(("family", "cap"), [("fact_eiffel", 20), ("fact_legal", 20), ("whois_snippet", 10)]) @@ -180,7 +192,7 @@ async def test_budget_coverage_and_full_resolution_async(self) -> None: assert {config._coverage_key(group) for group in flat} == set(config.coverage_keys) full = await config.get_attack_seed_groups_async(apply_sampling=False) assert len(full) == 12 - config.max_dataset_size = None + config.max_dataset_size = "all" assert len(await config.get_attack_seed_groups_async()) == 12 for group in full: assert group.prompts[0].source == "https://example.test/context" @@ -193,8 +205,8 @@ async def test_budget_coverage_and_full_resolution_async(self) -> None: ) async def test_runtime_budget_is_validated_without_sampling_async(self, cap: int, message: str) -> None: config = _config(families=["whois", "resume"]) - config.max_dataset_size = cap with pytest.raises(DatasetConstraintError, match=message): + config.max_total = cap await config.get_attack_seed_groups_async(apply_sampling=False) async def test_filters_validators_and_config_identity_async(self, seeded_memory_async: MemoryInterface) -> None: @@ -213,7 +225,8 @@ async def test_filters_validators_and_config_identity_async(self, seeded_memory_ await _initialize_async( scenario, dataset_config=config, scenario_techniques=[LatentInjectionTechnique.Bare] ) - assert scenario._dataset_config is config + assert scenario._dataset_config is not config + assert type(scenario._dataset_config) is type(config) assert config.families == ["whois"] assert len(seen) == 1 assert len(seen[0].seeds) == 24 diff --git a/tests/unit/scenario/garak/test_package_hallucination.py b/tests/unit/scenario/garak/test_package_hallucination.py index d4a82fa670..d5e1d4b2e4 100644 --- a/tests/unit/scenario/garak/test_package_hallucination.py +++ b/tests/unit/scenario/garak/test_package_hallucination.py @@ -7,9 +7,10 @@ import pytest +from pyrit.datasets.seed_datasets.seed_dataset_provider import SeedDatasetProvider from pyrit.executor.attack import PromptSendingAttack from pyrit.memory import MemoryInterface -from pyrit.models import AttackSeedGroup, ComponentIdentifier, SeedObjective, SeedPrompt +from pyrit.models import AttackSeedGroup, ComponentIdentifier, SeedDataset, SeedObjective, SeedPrompt from pyrit.prompt_target import PromptTarget from pyrit.scenario.core.dataset_configuration import DatasetConfiguration from pyrit.scenario.core.scenario import BaselineAttackPolicy @@ -60,6 +61,7 @@ def _get_seeds(*, dataset_name): return [MagicMock(value=value) for value in packages_by_dataset.get(dataset_name, [])] memory = MagicMock(spec=MemoryInterface) + memory.get_seed_dataset_names_async = AsyncMock(side_effect=lambda: list(packages_by_dataset)) memory.get_seeds_async = AsyncMock(side_effect=_get_seeds) memory.packages_by_dataset = packages_by_dataset return memory @@ -234,15 +236,21 @@ async def test_non_default_registry_is_fetched_lazily( ): fake_registry_memory.packages_by_dataset.pop(dataset_name) - async def _fetch_dataset_async(*, dataset_name: str) -> None: + async def store_dataset(**_: object) -> None: fake_registry_memory.packages_by_dataset[dataset_name] = packages - fetch_mock = AsyncMock(side_effect=_fetch_dataset_async) - with patch.object(DatasetConfiguration, "_fetch_dataset_async", new=fetch_mock): + provider = MagicMock(spec=SeedDatasetProvider) + provider.fetch_dataset_async.return_value = SeedDataset( + seeds=[SeedPrompt(value=value, dataset_name=dataset_name) for value in packages], + dataset_name=dataset_name, + ) + fake_registry_memory.add_seed_datasets_to_memory_async.side_effect = store_dataset + with patch.object(SeedDatasetProvider, "get_providers_by_name_async", return_value={dataset_name: provider}): scenario = PackageHallucination() await self._initialize(scenario, mock_objective_target, [technique], fake_registry_memory) - fetch_mock.assert_awaited_once_with(dataset_name=dataset_name) + provider.fetch_dataset_async.assert_awaited_once() + fake_registry_memory.add_seed_datasets_to_memory_async.assert_awaited_once() scorer = scenario._atomic_attacks[0].attack_technique.attack._objective_scorer assert scorer._ecosystem is ecosystem @@ -278,7 +286,7 @@ async def test_missing_corpus_raises(self, mock_objective_target): "pyrit.scenario.core.dataset_configuration.CentralMemory.get_memory_instance", return_value=empty_memory, ), - patch.object(DatasetConfiguration, "_fetch_dataset_async", new_callable=AsyncMock), + patch.object(DatasetConfiguration, "prepare_async", new_callable=AsyncMock), ): with pytest.raises(ValueError): scenario.set_params_from_args( @@ -309,7 +317,7 @@ def _get_seeds(*, dataset_name): "pyrit.scenario.core.dataset_configuration.CentralMemory.get_memory_instance", return_value=corpus_only, ), - patch.object(DatasetConfiguration, "_fetch_dataset_async", new_callable=AsyncMock), + patch.object(DatasetConfiguration, "prepare_async", new_callable=AsyncMock), ): with pytest.raises(ValueError): scenario.set_params_from_args( diff --git a/tests/unit/scenario/garak/test_prompt_inject.py b/tests/unit/scenario/garak/test_prompt_inject.py index 0c7bdbc95e..252d52b394 100644 --- a/tests/unit/scenario/garak/test_prompt_inject.py +++ b/tests/unit/scenario/garak/test_prompt_inject.py @@ -176,7 +176,7 @@ async def test_uncapped_configuration_uses_complete_matrix(self, mock_objective_ target=mock_objective_target, dataset_config=PromptInjectDatasetConfiguration( dataset_names=PromptInject.required_datasets(), - max_dataset_size=None, + max_dataset_size="all", ), ) @@ -236,7 +236,7 @@ async def test_converters_preserve_rendered_prompts_async( target=mock_objective_target, goal_texts=[goal], dataset_config=PromptInjectDatasetConfiguration( - dataset_names=PromptInject.required_datasets(), max_dataset_size=None + dataset_names=PromptInject.required_datasets(), max_dataset_size="all" ), ) @@ -403,11 +403,14 @@ async def test_auto_fetch_false_is_preserved_async(self, mock_objective_target: scenario = PromptInject() config = PromptInjectDatasetConfiguration(dataset_names=PromptInject.required_datasets(), auto_fetch=False) - with patch.object(config, "_fetch_dataset_async") as fetch: - with pytest.raises(DatasetConstraintError, match="auto_fetch is disabled"): + with patch( + "pyrit.datasets.seed_datasets.seed_dataset_provider.SeedDatasetProvider.get_providers_by_name_async" + ) as fetch: + with pytest.raises(DatasetConstraintError, match="fetch is 'never'"): await _initialize_async(scenario, target=mock_objective_target, dataset_config=config) fetch.assert_not_called() - assert scenario._dataset_config is config + assert scenario._dataset_config is not config + assert type(scenario._dataset_config) is type(config) async def test_unsupported_dataset_configuration_type_raises(self, mock_objective_target: PromptTarget) -> None: scenario = PromptInject() @@ -455,15 +458,24 @@ async def test_inline_dataset_is_rejected(self, mock_objective_target: PromptTar class TestPromptInjectDatasetSampling: @pytest.mark.parametrize( ("kwargs", "expected"), - [({}, 12), ({"max_dataset_size": None}, 210), ({"max_dataset_size": 6}, 6)], + [ + ({}, 12), + ({"max_dataset_size": None}, 12), + ({"max_total": "default"}, 12), + ({"max_total": ""}, 12), + ({"max_dataset_size": "all"}, 210), + ({"max_total": "all"}, 210), + ({"max_dataset_size": 6}, 6), + ], ) async def test_configuration_default_covers_custom_goals_async( - self, *, kwargs: dict[str, int | None], expected: int + self, *, kwargs: dict[str, int | str | None], expected: int ) -> None: goals = [f"goal {index}" for index in range(6)] config = PromptInjectDatasetConfiguration( dataset_names=PromptInject.required_datasets(), goal_texts=goals, **kwargs ) + await config.prepare_async() groups = await config.get_attack_seed_groups_async() assert len(groups) == expected assert {group.objective.metadata["goal_text"] for group in groups} == set(goals) @@ -474,6 +486,7 @@ async def test_both_resolvers_preserve_goal_coverage_async(self, grouped: bool) config = PromptInjectDatasetConfiguration( dataset_names=PromptInject.required_datasets(), goal_texts=goals, max_dataset_size=3 ) + await config.prepare_async() with patch("pyrit.scenario.scenarios.garak._prompt_injection.random", random.Random(0)): if grouped: by_dataset = await config.get_attack_groups_by_dataset_async() diff --git a/tests/unit/scenario/garak/test_system_prompt_extraction.py b/tests/unit/scenario/garak/test_system_prompt_extraction.py index 146562c345..147bbdfddd 100644 --- a/tests/unit/scenario/garak/test_system_prompt_extraction.py +++ b/tests/unit/scenario/garak/test_system_prompt_extraction.py @@ -61,6 +61,11 @@ def mock_objective_scorer(): @pytest.mark.usefixtures("patch_central_database") class TestSystemPromptExtractionInitialization: + @pytest.mark.parametrize("cap", [0, -1, True, 1.5, "unlimited"]) + def test_invalid_prompt_cap_raises(self, cap: object) -> None: + with pytest.raises(ValueError, match="positive integer"): + SystemPromptExtraction(prompt_cap=cap) + def test_no_arg_construction_for_registry(self): scenario = SystemPromptExtraction() assert scenario.name == "SystemPromptExtraction" @@ -100,6 +105,7 @@ def test_technique_all_expands_to_every_category(self): class TestSystemPromptExtractionAtomicAttacks: async def _init(self, scenario, mock_objective_target, techniques=None): with ( + patch.object(scenario._dataset_config, "prepare_async", new_callable=AsyncMock), patch.object(scenario._dataset_config, "_collect_named_seeds_async", new_callable=AsyncMock), patch.object(SystemPromptExtraction, "_load_system_prompts_async", return_value=list(SYSTEM_PROMPTS)), patch.object( @@ -177,8 +183,8 @@ async def test_prompt_cap_limits_total_seed_groups(self, mock_objective_target, total = sum(len(a.seed_groups) for a in scenario._atomic_attacks) assert total == 5 - async def test_prompt_cap_none_runs_every_combination(self, mock_objective_target, mock_objective_scorer): - scenario = SystemPromptExtraction(objective_scorer=mock_objective_scorer, prompt_cap=None) + async def test_prompt_cap_all_runs_every_combination(self, mock_objective_target, mock_objective_scorer): + scenario = SystemPromptExtraction(objective_scorer=mock_objective_scorer, prompt_cap="all") await self._init(scenario, mock_objective_target) total = sum(len(a.seed_groups) for a in scenario._atomic_attacks) diff --git a/tests/unit/scenario/scenarios/adaptive/test_text_adaptive.py b/tests/unit/scenario/scenarios/adaptive/test_text_adaptive.py index 81fa4adcf5..73a3928e9b 100644 --- a/tests/unit/scenario/scenarios/adaptive/test_text_adaptive.py +++ b/tests/unit/scenario/scenarios/adaptive/test_text_adaptive.py @@ -15,7 +15,7 @@ from pyrit.models.identifiers import ComponentIdentifier from pyrit.prompt_target import PromptTarget from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry -from pyrit.scenario.core.dataset_configuration import CompoundDatasetAttackConfiguration +from pyrit.scenario.core.dataset_configuration import DatasetAttackConfiguration from pyrit.scenario.core.scenario import BaselineAttackPolicy from pyrit.scenario.scenarios.adaptive.dispatcher import AdaptiveTechniqueDispatcher from pyrit.scenario.scenarios.adaptive.text_adaptive import TextAdaptive @@ -128,8 +128,9 @@ def test_baseline_enabled(self): def test_default_dataset_config(self): config = TextAdaptive.default_dataset_config() - assert isinstance(config, CompoundDatasetAttackConfiguration) - assert all(child.max_dataset_size == 4 for child in config._configurations) + assert isinstance(config, DatasetAttackConfiguration) + assert config.max_per_dataset == 4 + assert config.max_total == "all" assert config.dataset_names == TextAdaptive.required_datasets() def test_required_datasets_non_empty(self): @@ -216,7 +217,7 @@ async def _build_scenario_and_attacks( **scenario_kwargs, ): with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=seed_groups, @@ -271,7 +272,7 @@ def _spy_init(self, *args, **kwargs): "hate": [_make_seed_group(value="obj-h1", harm_categories=["hate"])], } with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -323,7 +324,7 @@ async def test_display_group_is_dataset_name(self, mock_objective_target, mock_o async def test_no_usable_techniques_raises(self, mock_objective_target, mock_objective_scorer): groups = {"violence": [_make_seed_group(value="obj")]} with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -351,7 +352,7 @@ async def test_techniques_with_seed_technique_are_kept(self, mock_objective_targ with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -398,7 +399,7 @@ async def test_incompatible_seed_technique_is_filtered_per_objective( # Only the plain factory (no seed_technique) is compatible. with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -452,7 +453,7 @@ def _selective_compat(self_group, *, technique): with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -498,7 +499,7 @@ class NarrowScoringConfig(AttackScoringConfig): narrow_factory = _make_fake_factory(scoring_config_type=NarrowScoringConfig) with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -548,7 +549,7 @@ def __init__(self, *, objective_scorer): with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -594,7 +595,7 @@ async def test_factory_create_failure_skips_technique(self, mock_objective_targe with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -631,7 +632,7 @@ async def test_all_factories_failing_raises_with_reason(self, mock_objective_tar with ( patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -661,7 +662,7 @@ class TestTextAdaptiveBaselinePolicy: async def test_initialize_async_accepts_explicit_baseline(self, mock_objective_target, mock_objective_scorer): groups = {"violence": [_make_seed_group(value="obj", harm_categories=["violence"])]} with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, @@ -683,7 +684,7 @@ async def test_baseline_emitted_at_index_zero_by_default(self, mock_objective_ta """ groups = {"violence": [_make_seed_group(value="obj", harm_categories=["violence"])]} with patch.object( - CompoundDatasetAttackConfiguration, + DatasetAttackConfiguration, "get_attack_groups_by_dataset_async", new_callable=AsyncMock, return_value=groups, diff --git a/tests/unit/scenario/test_default_run_size_estimates.py b/tests/unit/scenario/test_default_run_size_estimates.py index aa8392150b..7253039cc5 100644 --- a/tests/unit/scenario/test_default_run_size_estimates.py +++ b/tests/unit/scenario/test_default_run_size_estimates.py @@ -41,7 +41,7 @@ from pyrit.scenario.scenarios.airt.psychosocial import Psychosocial from pyrit.scenario.scenarios.benchmark.adversarial import AdversarialBenchmark from pyrit.scenario.scenarios.foundry.red_team_agent import FoundryComposite, FoundryTechnique, RedTeamAgent -from pyrit.scenario.scenarios.garak.api_key import ApiKey +from pyrit.scenario.scenarios.garak.api_key import ApiKey, ApiKeyDatasetConfiguration from pyrit.scenario.scenarios.garak.encoding import Encoding from pyrit.scenario.scenarios.garak.exploitation import Exploitation from pyrit.scenario.scenarios.garak.figstep import FigStep @@ -184,8 +184,6 @@ async def test_default_estimate_uses_five_without_population_or_persistence_asyn ), (Exploitation, {}, {"prompt_cap": 0}, "prompt_cap must be greater than zero"), (Exploitation, {}, {"prompt_cap": -1}, "prompt_cap must be greater than zero"), - (SystemPromptExtraction, {"prompt_cap": 0}, {}, "prompt_cap must be greater than zero"), - (SystemPromptExtraction, {"prompt_cap": -1}, {}, "prompt_cap must be greater than zero"), (SystemPromptExtraction, {"system_prompt_subsample": 0}, {}, "system_prompt_subsample"), (PackageHallucination, {"max_prompts_per_language": 0}, {}, "max_prompts_per_language"), (PackageHallucination, {"max_prompts_per_language": -1}, {}, "max_prompts_per_language"), @@ -223,8 +221,22 @@ async def test_configuration_only_checks_are_shared_before_dataset_reads_async( @pytest.mark.usefixtures("patch_central_database") -@pytest.mark.parametrize("cap", [3, 12, None]) -async def test_prompt_inject_valid_coverage_caps_still_have_configuration_only_previews_async(cap: int | None) -> None: +@pytest.mark.parametrize( + "configuration_class", + [ApiKeyDatasetConfiguration, PromptInjectDatasetConfiguration, LatentInjectionDatasetConfiguration], +) +def test_ingredient_configurations_accept_explicit_sampling_scope( + configuration_class: type[DatasetAttackConfiguration], +) -> None: + config = configuration_class(sampling_scope="total_only") + assert config.max_per_dataset == "all" + with pytest.raises(ValueError, match="requires total_only"): + configuration_class(sampling_scope="per_dataset") + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("cap", [3, 12, "all"]) +async def test_prompt_inject_valid_coverage_caps_still_have_configuration_only_previews_async(cap: int | str) -> None: scenario = PromptInject() scenario.set_params_from_args( args={ @@ -234,7 +246,7 @@ async def test_prompt_inject_valid_coverage_caps_still_have_configuration_only_p } ) estimate = await scenario.get_run_size_estimate_async() - if cap is None: + if cap == "all": assert estimate.status is ScenarioRunSizeEstimateStatus.Unavailable else: assert estimate.estimated_attack_count == cap * 5 @@ -275,7 +287,7 @@ async def test_unavailable_formula_preserves_known_budget_async() -> None: estimate = await scenario.get_run_size_estimate_async() assert estimate.status is ScenarioRunSizeEstimateStatus.Unavailable assert estimate.dataset_size == BoundedDatasetSize(value=5) - assert estimate.dataset_limit.value == 5 + assert estimate.dataset_limit.value is None assert estimate.note == "No formula." @@ -287,7 +299,9 @@ async def test_configured_estimate_uses_selected_techniques_and_limit_async(*, b args={ "scenario_techniques": [_TwoTechniqueDefault.ONE], "include_baseline": baseline, - "dataset_config": DatasetAttackConfiguration(dataset_names=["also-missing"], max_dataset_size=7), + "dataset_config": DatasetAttackConfiguration( + dataset_names=["also-missing"], max_per_dataset="all", max_total=7 + ), } ) estimate = await scenario.get_run_size_estimate_async() @@ -303,7 +317,9 @@ async def test_estimate_expands_aggregate_and_applies_combined_cap_once_async() args={ "scenario_techniques": [_TwoTechniqueDefault.ALL], "include_baseline": False, - "dataset_config": DatasetAttackConfiguration(dataset_names=["one", "two"], max_dataset_size=3), + "dataset_config": DatasetAttackConfiguration( + dataset_names=["one", "two"], max_per_dataset="all", max_total=3 + ), } ) estimate = await scenario.get_run_size_estimate_async() @@ -330,7 +346,11 @@ async def test_estimate_combines_independent_child_limits_async() -> None: async def test_unlimited_estimate_does_not_load_data_or_invent_a_count_async() -> None: scenario = _MatrixEstimateScenario() scenario.set_params_from_args( - args={"dataset_config": DatasetAttackConfiguration(dataset_names=["missing"], max_dataset_size=None)} + args={ + "dataset_config": DatasetAttackConfiguration( + dataset_names=["missing"], max_per_dataset="all", max_total="all" + ) + } ) estimate = await scenario.get_run_size_estimate_async() assert estimate.status is ScenarioRunSizeEstimateStatus.Unavailable @@ -517,7 +537,9 @@ async def test_web_injection_capped_techniques_use_generation_limits_async( WebInjectionTechnique.TaskXSS, ], "include_baseline": baseline, - "dataset_config": DatasetAttackConfiguration(dataset_names=["missing"], max_dataset_size=dataset_limit), + "dataset_config": DatasetAttackConfiguration( + dataset_names=["missing"], max_per_dataset="all", max_total=dataset_limit + ), } ) estimate = await scenario.get_run_size_estimate_async() @@ -539,7 +561,9 @@ async def test_package_hallucination_uses_per_language_generation_cap_async( scenario.set_params_from_args( args={ "scenario_techniques": [technique], - "dataset_config": DatasetAttackConfiguration(dataset_names=["missing"], max_dataset_size=dataset_limit), + "dataset_config": DatasetAttackConfiguration( + dataset_names=["missing"], max_per_dataset="all", max_total=dataset_limit + ), } ) estimate = await scenario.get_run_size_estimate_async() @@ -552,29 +576,32 @@ async def test_package_hallucination_uses_per_language_generation_cap_async( @pytest.mark.usefixtures("patch_central_database") -@pytest.mark.parametrize("prompt_cap", [256, 7, None]) +@pytest.mark.parametrize("prompt_cap", [256, 7, None, "", "default", "all"]) @pytest.mark.parametrize( "technique", [SystemPromptExtractionTechnique.ALL, SystemPromptExtractionTechnique.DirectRequests] ) @pytest.mark.parametrize("dataset_limit", [1, None]) async def test_system_prompt_extraction_uses_one_shared_generation_cap_async( - *, prompt_cap: int | None, technique: SystemPromptExtractionTechnique, dataset_limit: int | None + *, prompt_cap: int | str | None, technique: SystemPromptExtractionTechnique, dataset_limit: int | None ) -> None: scenario = SystemPromptExtraction(objective_scorer=_scorer(), prompt_cap=prompt_cap) scenario.set_params_from_args( args={ "scenario_techniques": [technique], - "dataset_config": DatasetAttackConfiguration(dataset_names=["missing"], max_dataset_size=dataset_limit), + "dataset_config": DatasetAttackConfiguration( + dataset_names=["missing"], max_per_dataset="all", max_total=dataset_limit + ), } ) estimate = await scenario.get_run_size_estimate_async() - assert estimate.estimated_attack_count == prompt_cap - assert estimate.dataset_size == scenario_dataset_size_from_limit(prompt_cap) - if prompt_cap is None: + expected = 256 if prompt_cap in (None, "", "default") else prompt_cap + assert estimate.estimated_attack_count == (None if expected == "all" else expected) + assert estimate.dataset_size == scenario_dataset_size_from_limit(expected) + if expected == "all": assert estimate.status is ScenarioRunSizeEstimateStatus.Unavailable else: assert estimate.status is ScenarioRunSizeEstimateStatus.Approximate - assert estimate.effective_parameters == {"prompt_cap": prompt_cap} + assert estimate.effective_parameters == {"prompt_cap": expected} @pytest.mark.usefixtures("patch_central_database") @@ -583,7 +610,8 @@ async def test_psychosocial_keeps_per_harm_limits_and_baselines_async() -> None: estimate = await scenario.get_default_run_size_estimate_async() assert estimate.estimated_attack_count == 40 assert estimate.dataset_size == scenario_dataset_size_from_limit(10) - assert [component.count for component in estimate.components] == [15, 5, 15, 5] + assert [component.count for component in estimate.components] == [30, 10] + assert [dataset.configured_caps[0].count for dataset in estimate.datasets] == [5, 5] @pytest.mark.usefixtures("patch_central_database") diff --git a/tests/unit/scenario/test_generated_objectives.py b/tests/unit/scenario/test_generated_objectives.py index a06c3c43ce..e2b3fd51e3 100644 --- a/tests/unit/scenario/test_generated_objectives.py +++ b/tests/unit/scenario/test_generated_objectives.py @@ -54,7 +54,9 @@ async def respond_async(*, normalized_conversation: list[Message]) -> list[Messa stored = await sqlite_instance.get_seeds_async(dataset_name=provider.dataset_name, origin=SeedOrigin.GENERATED) assert len(stored) == 10 - config = DatasetAttackConfiguration(dataset_names=[provider.dataset_name], max_dataset_size=None, auto_fetch=False) + config = DatasetAttackConfiguration( + dataset_names=[provider.dataset_name], max_per_dataset="all", max_total="all", auto_fetch=False + ) target = MockPromptTarget() scenario = RapidResponse(objective_scorer=SubStringScorer(substring="default")) scenario.set_params_from_args( From b69cfdc94e750598494bc5aac56f10c5ef9c444c Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 9 Oct 2026 18:56:57 -0700 Subject: [PATCH 2/2] Fix scenario dataset review guidance and examples Keep the deprecated total alias exact, correct Psychosocial limit advice, and migrate paired documentation examples to explicit dataset sources. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../instructions/scenarios.instructions.md | 5 ++++ doc/code/datasets/6_generated_datasets.ipynb | 4 +-- doc/code/datasets/6_generated_datasets.py | 4 +-- doc/code/scenarios/0_scenarios.ipynb | 1 + doc/code/scenarios/0_scenarios.py | 1 + .../1_common_scenario_parameters.ipynb | 10 +++---- .../scenarios/1_common_scenario_parameters.py | 10 +++---- doc/code/scenarios/3_adaptive_scenarios.ipynb | 10 +++---- doc/code/scenarios/3_adaptive_scenarios.py | 10 +++---- doc/scanner/1_pyrit_scan.ipynb | 6 ++-- doc/scanner/1_pyrit_scan.py | 6 ++-- doc/scanner/adaptive.ipynb | 4 +-- doc/scanner/adaptive.py | 4 +-- doc/scanner/airt.ipynb | 16 +++++----- doc/scanner/airt.py | 16 +++++----- doc/scanner/benchmark.ipynb | 4 +-- doc/scanner/benchmark.py | 4 +-- doc/scanner/foundry.ipynb | 4 +-- doc/scanner/foundry.py | 4 +-- doc/scanner/garak.ipynb | 30 +++++++++++-------- doc/scanner/garak.py | 30 +++++++++++-------- pyrit/scenario/scenarios/airt/psychosocial.py | 3 +- tests/unit/scenario/airt/test_psychosocial.py | 5 +++- 23 files changed, 105 insertions(+), 86 deletions(-) diff --git a/.github/instructions/scenarios.instructions.md b/.github/instructions/scenarios.instructions.md index ad3d19a745..bef7214bd7 100644 --- a/.github/instructions/scenarios.instructions.md +++ b/.github/instructions/scenarios.instructions.md @@ -167,6 +167,11 @@ Options: Reads use `_collect_seeds_for_dataset_async()` and must not fetch. - `dataset_names`, `max_dataset_size`, `auto_fetch`, and `per_dataset()` are deprecated. +`max_dataset_size` is an exact alias for `max_total`; it does not disable source caps. +To keep the old total-only selection of up to 10 groups, use +`max_per_dataset="all", max_total=10`. To remove both caps, set both to `"all"`; +`None` now uses the default, not unlimited selection. + ## Technique Enum Technique members should represent **attack techniques** — the *how* of an attack (e.g., prompt sending, role play, TAP). Datasets control *what* is tested (e.g., harm categories, compliance topics). Avoid mixing dataset/category selection into the technique enum; use `DatasetConfiguration` and `--dataset-names` for that axis. diff --git a/doc/code/datasets/6_generated_datasets.ipynb b/doc/code/datasets/6_generated_datasets.ipynb index 0210af9941..334f72bdf4 100644 --- a/doc/code/datasets/6_generated_datasets.ipynb +++ b/doc/code/datasets/6_generated_datasets.ipynb @@ -240,11 +240,11 @@ ], "source": [ "from pyrit.output import output_scenario_async, output_scenario_attacks_async\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetFetchPolicy, DatasetSource\n", "from pyrit.scenario.airt import RapidResponse, RapidResponseTechnique\n", "from pyrit.score import SelfAskTrueFalseScorer\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[dataset_name], auto_fetch=False)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=dataset_name)], fetch=DatasetFetchPolicy.NEVER)\n", "scenario = RapidResponse(objective_scorer=SelfAskTrueFalseScorer(chat_target=OpenAIChatTarget()))\n", "scenario.set_params_from_args(\n", " args={\n", diff --git a/doc/code/datasets/6_generated_datasets.py b/doc/code/datasets/6_generated_datasets.py index 44812af346..c0d228853b 100644 --- a/doc/code/datasets/6_generated_datasets.py +++ b/doc/code/datasets/6_generated_datasets.py @@ -76,11 +76,11 @@ # %% from pyrit.output import output_scenario_async, output_scenario_attacks_async -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetFetchPolicy, DatasetSource from pyrit.scenario.airt import RapidResponse, RapidResponseTechnique from pyrit.score import SelfAskTrueFalseScorer -dataset_config = DatasetAttackConfiguration(dataset_names=[dataset_name], auto_fetch=False) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=dataset_name)], fetch=DatasetFetchPolicy.NEVER) scenario = RapidResponse(objective_scorer=SelfAskTrueFalseScorer(chat_target=OpenAIChatTarget())) scenario.set_params_from_args( args={ diff --git a/doc/code/scenarios/0_scenarios.ipynb b/doc/code/scenarios/0_scenarios.ipynb index 711abd6756..b088db7741 100644 --- a/doc/code/scenarios/0_scenarios.ipynb +++ b/doc/code/scenarios/0_scenarios.ipynb @@ -68,6 +68,7 @@ " - Returns a `DatasetAttackConfiguration` with named sources (e.g., `sources=[DatasetSource(name=\"my_dataset\")]`)\n", " - Users can override this at runtime via `--dataset-names` in the CLI or by passing a custom `dataset_config` programmatically\n", " - Named sources select at most 5 attack groups per dataset by default; `max_total` caps the union. For any limit, omitted/`None`/empty/`\"default\"` uses the default; `\"all\"` removes that limit.\n", + " - The deprecated `max_dataset_size` is an exact alias for `max_total`, not a way to disable source caps. To keep the old total-only selection, use `max_per_dataset=\"all\", max_total=10`; to remove both caps, set both to `\"all\"`. `None` now uses the default, not unlimited selection.\n", " - New runs prepare missing registered datasets once; reads, estimates, and resume never fetch\n", "\n", "4. **Constructor**: Use `@apply_defaults` decorator and call `super().__init__()` with scenario metadata:\n", diff --git a/doc/code/scenarios/0_scenarios.py b/doc/code/scenarios/0_scenarios.py index 2cf481886b..3a951368a4 100644 --- a/doc/code/scenarios/0_scenarios.py +++ b/doc/code/scenarios/0_scenarios.py @@ -70,6 +70,7 @@ # - Returns a `DatasetAttackConfiguration` with named sources (e.g., `sources=[DatasetSource(name="my_dataset")]`) # - Users can override this at runtime via `--dataset-names` in the CLI or by passing a custom `dataset_config` programmatically # - Named sources select at most 5 attack groups per dataset by default; `max_total` caps the union. For any limit, omitted/`None`/empty/`"default"` uses the default; `"all"` removes that limit. +# - The deprecated `max_dataset_size` is an exact alias for `max_total`, not a way to disable source caps. To keep the old total-only selection, use `max_per_dataset="all", max_total=10`; to remove both caps, set both to `"all"`. `None` now uses the default, not unlimited selection. # - New runs prepare missing registered datasets once; reads, estimates, and resume never fetch # # 4. **Constructor**: Use `@apply_defaults` decorator and call `super().__init__()` with scenario metadata: diff --git a/doc/code/scenarios/1_common_scenario_parameters.ipynb b/doc/code/scenarios/1_common_scenario_parameters.ipynb index af9ff31ffc..e6081ccad4 100644 --- a/doc/code/scenarios/1_common_scenario_parameters.ipynb +++ b/doc/code/scenarios/1_common_scenario_parameters.ipynb @@ -75,7 +75,7 @@ "## Dataset Configuration\n", "\n", "`DatasetAttackConfiguration` controls which prompts (objectives) are sent to the target.\n", - "The simplest approach uses `dataset_names` to load datasets by name from memory.\n", + "Use `sources` to select named datasets from memory.\n", "By default, `RedTeamAgent` loads four random objectives from HarmBench [@mazeika2024harmbench]." ] }, @@ -86,9 +86,9 @@ "metadata": {}, "outputs": [], "source": [ - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"harmbench\"], max_dataset_size=2)" + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")], max_total=2)" ] }, { @@ -129,8 +129,8 @@ " if all(seed.value.isascii() and seed.value.isprintable() for seed in group.seeds)\n", "]\n", "\n", - "# Pass explicit seed_groups instead of dataset_names\n", - "dataset_config = DatasetAttackConfiguration(seed_groups=seed_groups, max_dataset_size=2)" + "# Pass explicit seed_groups instead of named sources\n", + "dataset_config = DatasetAttackConfiguration(seed_groups=seed_groups, max_total=2)" ] }, { diff --git a/doc/code/scenarios/1_common_scenario_parameters.py b/doc/code/scenarios/1_common_scenario_parameters.py index fff1d1dba2..354df1410e 100644 --- a/doc/code/scenarios/1_common_scenario_parameters.py +++ b/doc/code/scenarios/1_common_scenario_parameters.py @@ -42,13 +42,13 @@ # ## Dataset Configuration # # `DatasetAttackConfiguration` controls which prompts (objectives) are sent to the target. -# The simplest approach uses `dataset_names` to load datasets by name from memory. +# Use `sources` to select named datasets from memory. # By default, `RedTeamAgent` loads four random objectives from HarmBench [@mazeika2024harmbench]. # %% -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource -dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=2) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")], max_total=2) # %% [markdown] # For more control, use `SeedDatasetProvider` to fetch datasets and pass explicit `seed_groups`. @@ -69,8 +69,8 @@ if all(seed.value.isascii() and seed.value.isprintable() for seed in group.seeds) ] -# Pass explicit seed_groups instead of dataset_names -dataset_config = DatasetAttackConfiguration(seed_groups=seed_groups, max_dataset_size=2) +# Pass explicit seed_groups instead of named sources +dataset_config = DatasetAttackConfiguration(seed_groups=seed_groups, max_total=2) # %% [markdown] # ## Technique Selection and Composition diff --git a/doc/code/scenarios/3_adaptive_scenarios.ipynb b/doc/code/scenarios/3_adaptive_scenarios.ipynb index 763bdcc840..2c9b2bc2f5 100644 --- a/doc/code/scenarios/3_adaptive_scenarios.ipynb +++ b/doc/code/scenarios/3_adaptive_scenarios.ipynb @@ -113,7 +113,7 @@ "\n", "from pyrit.output.scenario_result.pretty import PrettyScenarioResultMemoryPrinter\n", "from pyrit.registry import TargetRegistry\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.scenario.scenarios.adaptive import TextAdaptive\n", "from pyrit.setup import initialize_from_config_async\n", "\n", @@ -452,8 +452,8 @@ " \"objective_target\": objective_target,\n", " \"scenario_techniques\": [technique_class(\"single_turn\")],\n", " \"dataset_config\": DatasetAttackConfiguration(\n", - " dataset_names=[\"airt_hate\", \"airt_violence\"],\n", - " max_dataset_size=4,\n", + " sources=[DatasetSource(name=\"airt_hate\"), DatasetSource(name=\"airt_violence\")],\n", + " max_total=4,\n", " ),\n", " }\n", ")\n", @@ -577,8 +577,8 @@ " \"objective_target\": objective_target,\n", " \"scenario_techniques\": [technique_class(\"single_turn\")],\n", " \"dataset_config\": DatasetAttackConfiguration(\n", - " dataset_names=[\"airt_hate\", \"airt_violence\"],\n", - " max_dataset_size=4,\n", + " sources=[DatasetSource(name=\"airt_hate\"), DatasetSource(name=\"airt_violence\")],\n", + " max_total=4,\n", " ),\n", " }\n", ")\n", diff --git a/doc/code/scenarios/3_adaptive_scenarios.py b/doc/code/scenarios/3_adaptive_scenarios.py index c4cc621f2e..247b24b786 100644 --- a/doc/code/scenarios/3_adaptive_scenarios.py +++ b/doc/code/scenarios/3_adaptive_scenarios.py @@ -48,7 +48,7 @@ from pyrit.output.scenario_result.pretty import PrettyScenarioResultMemoryPrinter from pyrit.registry import TargetRegistry -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.scenarios.adaptive import TextAdaptive from pyrit.setup import initialize_from_config_async @@ -102,8 +102,8 @@ "objective_target": objective_target, "scenario_techniques": [technique_class("single_turn")], "dataset_config": DatasetAttackConfiguration( - dataset_names=["airt_hate", "airt_violence"], - max_dataset_size=4, + sources=[DatasetSource(name="airt_hate"), DatasetSource(name="airt_violence")], + max_total=4, ), } ) @@ -132,8 +132,8 @@ "objective_target": objective_target, "scenario_techniques": [technique_class("single_turn")], "dataset_config": DatasetAttackConfiguration( - dataset_names=["airt_hate", "airt_violence"], - max_dataset_size=4, + sources=[DatasetSource(name="airt_hate"), DatasetSource(name="airt_violence")], + max_total=4, ), } ) diff --git a/doc/scanner/1_pyrit_scan.ipynb b/doc/scanner/1_pyrit_scan.ipynb index 751d149e0f..ee28f8b9e7 100644 --- a/doc/scanner/1_pyrit_scan.ipynb +++ b/doc/scanner/1_pyrit_scan.ipynb @@ -1287,7 +1287,7 @@ "\n", "from pyrit.common import apply_defaults\n", "from pyrit.prompt_target.openai.openai_chat_target import OpenAIChatTarget\n", - "from pyrit.scenario import DatasetAttackConfiguration, Scenario, ScenarioTechnique\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource, Scenario, ScenarioTechnique\n", "from pyrit.score import SelfAskRefusalScorer, TrueFalseInverterScorer\n", "from pyrit.setup import initialize_pyrit_async\n", "\n", @@ -1311,8 +1311,8 @@ " version=1,\n", " objective_scorer=TrueFalseInverterScorer(scorer=SelfAskRefusalScorer(chat_target=OpenAIChatTarget())),\n", " technique_class=MyCustomTechnique,\n", - " # DatasetAttackConfiguration selects at most 5 attack groups by default; set max_dataset_size to change it.\n", - " default_dataset_config=DatasetAttackConfiguration(dataset_names=[\"harmbench\"]),\n", + " # Named sources default to 5 groups each; use max_per_dataset and max_total to set caps.\n", + " default_dataset_config=DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")]),\n", " scenario_result_id=scenario_result_id,\n", " )\n", " # ... your scenario-specific initialization code\n", diff --git a/doc/scanner/1_pyrit_scan.py b/doc/scanner/1_pyrit_scan.py index 7e1d98d5c3..0aa18722aa 100644 --- a/doc/scanner/1_pyrit_scan.py +++ b/doc/scanner/1_pyrit_scan.py @@ -182,7 +182,7 @@ from pyrit.common import apply_defaults from pyrit.prompt_target.openai.openai_chat_target import OpenAIChatTarget -from pyrit.scenario import DatasetAttackConfiguration, Scenario, ScenarioTechnique +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource, Scenario, ScenarioTechnique from pyrit.score import SelfAskRefusalScorer, TrueFalseInverterScorer from pyrit.setup import initialize_pyrit_async @@ -206,8 +206,8 @@ def __init__(self, *, scenario_result_id=None, **kwargs): version=1, objective_scorer=TrueFalseInverterScorer(scorer=SelfAskRefusalScorer(chat_target=OpenAIChatTarget())), technique_class=MyCustomTechnique, - # DatasetAttackConfiguration selects at most 5 attack groups by default; set max_dataset_size to change it. - default_dataset_config=DatasetAttackConfiguration(dataset_names=["harmbench"]), + # Named sources default to 5 groups each; use max_per_dataset and max_total to set caps. + default_dataset_config=DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")]), scenario_result_id=scenario_result_id, ) # ... your scenario-specific initialization code diff --git a/doc/scanner/adaptive.ipynb b/doc/scanner/adaptive.ipynb index 61a640dc3d..a4bc5849a4 100644 --- a/doc/scanner/adaptive.ipynb +++ b/doc/scanner/adaptive.ipynb @@ -107,7 +107,7 @@ "\n", "from pyrit.output import output_scenario_async\n", "from pyrit.registry import TargetRegistry\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.scenario.adaptive import TextAdaptive\n", "from pyrit.setup import initialize_from_config_async\n", "\n", @@ -115,7 +115,7 @@ "\n", "objective_target = TargetRegistry.get_registry_singleton().instances.get(\"openai_chat\")\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_hate\"], max_dataset_size=2)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_hate\")], max_total=2)\n", "\n", "scenario = TextAdaptive()\n", "scenario.set_params_from_args( # type: ignore\n", diff --git a/doc/scanner/adaptive.py b/doc/scanner/adaptive.py index 68e8f386c9..b229e0e55d 100644 --- a/doc/scanner/adaptive.py +++ b/doc/scanner/adaptive.py @@ -35,7 +35,7 @@ from pyrit.output import output_scenario_async from pyrit.registry import TargetRegistry -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.adaptive import TextAdaptive from pyrit.setup import initialize_from_config_async @@ -43,7 +43,7 @@ objective_target = TargetRegistry.get_registry_singleton().instances.get("openai_chat") -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_hate"], max_dataset_size=2) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_hate")], max_total=2) scenario = TextAdaptive() scenario.set_params_from_args( # type: ignore diff --git a/doc/scanner/airt.ipynb b/doc/scanner/airt.ipynb index 77b63150c5..09cf12e75e 100644 --- a/doc/scanner/airt.ipynb +++ b/doc/scanner/airt.ipynb @@ -74,7 +74,7 @@ "source": [ "from pyrit.output import output_scenario_async\n", "from pyrit.prompt_target import OpenAIChatTarget\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.setup import IN_MEMORY, initialize_pyrit_async\n", "from pyrit.setup.initializers import ScorerInitializer, TargetInitializer, TechniqueInitializer\n", "\n", @@ -133,7 +133,7 @@ "source": [ "from pyrit.scenario.airt import RapidResponse, RapidResponseTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_hate\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_hate\")], max_total=1)\n", "\n", "scenario = RapidResponse()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -287,7 +287,7 @@ "\n", "# Minimal demo: a single sub-harm, one technique (the bare simulated-crescendo base), and one\n", "# objective. Omit `scenario_techniques` to run the DEFAULT converter sweep across the full dataset.\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_imminent_crisis\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_imminent_crisis\")], max_total=1)\n", "\n", "scenario = Psychosocial()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -444,7 +444,7 @@ "source": [ "from pyrit.scenario.airt import Cyber, CyberTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_malware\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_malware\")], max_total=1)\n", "\n", "scenario = Cyber()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -604,7 +604,7 @@ "source": [ "from pyrit.scenario.airt import Jailbreak, JailbreakTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"harmbench\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")], max_total=1)\n", "\n", "scenario = Jailbreak()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -760,7 +760,7 @@ "source": [ "from pyrit.scenario.airt import Multilingual\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"harmbench\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")], max_total=1)\n", "\n", "scenario = Multilingual()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -919,7 +919,7 @@ "source": [ "from pyrit.scenario.airt import Leakage, LeakageTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_leakage\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_leakage\")], max_total=1)\n", "\n", "scenario = Leakage()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -1081,7 +1081,7 @@ "source": [ "from pyrit.scenario.airt import Scam, ScamTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"airt_scams\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"airt_scams\")], max_total=1)\n", "\n", "scenario = Scam()\n", "scenario.set_params_from_args( # type: ignore\n", diff --git a/doc/scanner/airt.py b/doc/scanner/airt.py index de13a2e44d..cb804f14e7 100644 --- a/doc/scanner/airt.py +++ b/doc/scanner/airt.py @@ -21,7 +21,7 @@ # %% from pyrit.output import output_scenario_async from pyrit.prompt_target import OpenAIChatTarget -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.setup import IN_MEMORY, initialize_pyrit_async from pyrit.setup.initializers import ScorerInitializer, TargetInitializer, TechniqueInitializer @@ -52,7 +52,7 @@ # %% from pyrit.scenario.airt import RapidResponse, RapidResponseTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_hate"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_hate")], max_total=1) scenario = RapidResponse() scenario.set_params_from_args( # type: ignore @@ -94,7 +94,7 @@ # Minimal demo: a single sub-harm, one technique (the bare simulated-crescendo base), and one # objective. Omit `scenario_techniques` to run the DEFAULT converter sweep across the full dataset. -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_imminent_crisis"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_imminent_crisis")], max_total=1) scenario = Psychosocial() scenario.set_params_from_args( # type: ignore @@ -133,7 +133,7 @@ # %% from pyrit.scenario.airt import Cyber, CyberTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_malware"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_malware")], max_total=1) scenario = Cyber() scenario.set_params_from_args( # type: ignore @@ -181,7 +181,7 @@ # %% from pyrit.scenario.airt import Jailbreak, JailbreakTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")], max_total=1) scenario = Jailbreak() scenario.set_params_from_args( # type: ignore @@ -228,7 +228,7 @@ # %% from pyrit.scenario.airt import Multilingual -dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")], max_total=1) scenario = Multilingual() scenario.set_params_from_args( # type: ignore @@ -281,7 +281,7 @@ # %% from pyrit.scenario.airt import Leakage, LeakageTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_leakage"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_leakage")], max_total=1) scenario = Leakage() scenario.set_params_from_args( # type: ignore @@ -318,7 +318,7 @@ # %% from pyrit.scenario.airt import Scam, ScamTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["airt_scams"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="airt_scams")], max_total=1) scenario = Scam() scenario.set_params_from_args( # type: ignore diff --git a/doc/scanner/benchmark.ipynb b/doc/scanner/benchmark.ipynb index 061cbab7ff..6abd678fe9 100644 --- a/doc/scanner/benchmark.ipynb +++ b/doc/scanner/benchmark.ipynb @@ -112,7 +112,7 @@ "source": [ "from pyrit.output import output_scenario_async\n", "from pyrit.prompt_target import OpenAIChatTarget\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.scenario.benchmark import AdversarialBenchmark\n", "from pyrit.setup import IN_MEMORY, initialize_pyrit_async\n", "from pyrit.setup.initializers import ScorerInitializer, TargetInitializer, TechniqueInitializer\n", @@ -149,7 +149,7 @@ "source": [ "from pyrit.scenario.benchmark import AdversarialBenchmarkTechnique\n", "\n", - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"harmbench\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")], max_total=1)\n", "\n", "scenario = AdversarialBenchmark()\n", "scenario.set_params_from_args(\n", diff --git a/doc/scanner/benchmark.py b/doc/scanner/benchmark.py index a9d1ba6f8b..8adddffbc1 100644 --- a/doc/scanner/benchmark.py +++ b/doc/scanner/benchmark.py @@ -56,7 +56,7 @@ # %% from pyrit.output import output_scenario_async from pyrit.prompt_target import OpenAIChatTarget -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.benchmark import AdversarialBenchmark from pyrit.setup import IN_MEMORY, initialize_pyrit_async from pyrit.setup.initializers import ScorerInitializer, TargetInitializer, TechniqueInitializer @@ -71,7 +71,7 @@ # %% from pyrit.scenario.benchmark import AdversarialBenchmarkTechnique -dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")], max_total=1) scenario = AdversarialBenchmark() scenario.set_params_from_args( diff --git a/doc/scanner/foundry.ipynb b/doc/scanner/foundry.ipynb index 06ad1c5190..78a5dbeaab 100644 --- a/doc/scanner/foundry.ipynb +++ b/doc/scanner/foundry.ipynb @@ -51,7 +51,7 @@ "\n", "from pyrit.output import output_scenario_async\n", "from pyrit.registry import TargetRegistry\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.scenario.foundry import FoundryTechnique, RedTeamAgent\n", "from pyrit.setup import initialize_from_config_async\n", "\n", @@ -116,7 +116,7 @@ } ], "source": [ - "dataset_config = DatasetAttackConfiguration(dataset_names=[\"harmbench\"], max_dataset_size=1)\n", + "dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"harmbench\")], max_total=1)\n", "\n", "scenario = RedTeamAgent()\n", "scenario.set_params_from_args( # type: ignore\n", diff --git a/doc/scanner/foundry.py b/doc/scanner/foundry.py index 74c21e403b..bbcc9b8bb6 100644 --- a/doc/scanner/foundry.py +++ b/doc/scanner/foundry.py @@ -24,7 +24,7 @@ from pyrit.output import output_scenario_async from pyrit.registry import TargetRegistry -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.foundry import FoundryTechnique, RedTeamAgent from pyrit.setup import initialize_from_config_async @@ -53,7 +53,7 @@ # | **Aggregates** | ALL, EASY, MODERATE, DIFFICULT | # %% -dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"], max_dataset_size=1) +dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="harmbench")], max_total=1) scenario = RedTeamAgent() scenario.set_params_from_args( # type: ignore diff --git a/doc/scanner/garak.ipynb b/doc/scanner/garak.ipynb index 2a1bdb7e66..e5e71b2f35 100644 --- a/doc/scanner/garak.ipynb +++ b/doc/scanner/garak.ipynb @@ -110,7 +110,7 @@ "\n", "from pyrit.output import output_scenario_async\n", "from pyrit.registry import TargetRegistry\n", - "from pyrit.scenario import DatasetAttackConfiguration\n", + "from pyrit.scenario import DatasetAttackConfiguration, DatasetSource\n", "from pyrit.scenario.garak import (\n", " ApiKey,\n", " ApiKeyDatasetConfiguration,\n", @@ -192,8 +192,8 @@ " \"scenario_techniques\": [PromptInjectTechnique.IgnorePrint],\n", " \"goal_texts\": [\"I hate humans\"],\n", " \"dataset_config\": PromptInjectDatasetConfiguration(\n", - " dataset_names=PromptInject.required_datasets(),\n", - " max_dataset_size=1,\n", + " sources=[DatasetSource(name=name) for name in PromptInject.required_datasets()],\n", + " max_total=1,\n", " ),\n", " }\n", ")\n", @@ -341,7 +341,7 @@ } ], "source": [ - "dataset_config = EncodingDatasetConfiguration(dataset_names=[\"garak_slur_terms_en\"], max_dataset_size=1)\n", + "dataset_config = EncodingDatasetConfiguration(sources=[DatasetSource(name=\"garak_slur_terms_en\")], max_total=1)\n", "\n", "scenario = Encoding()\n", "scenario.set_params_from_args( # type: ignore\n", @@ -581,7 +581,7 @@ } ], "source": [ - "figstep_dataset_config = DatasetAttackConfiguration(dataset_names=[\"figstep\"], max_dataset_size=1)\n", + "figstep_dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"figstep\")], max_total=1)\n", "\n", "figstep_scenario = FigStep()\n", "figstep_scenario.set_params_from_args( # type: ignore\n", @@ -968,7 +968,7 @@ "```\n", "\n", "**Available techniques:** `GetKey` and `CompleteKey`. `DEFAULT` and `ALL` both select the two\n", - "techniques. `max_dataset_size` samples across all selected technique populations, not per service.\n", + "techniques. `max_total` samples across all selected technique populations, not per service.\n", "The base scenario persists the sample for resume. Use `ApiKeyDatasetConfiguration` with\n", "`max_total=\"all\"` to run all 348 requests. Standard technique converter stacks are supported." ] @@ -1029,7 +1029,9 @@ " args={\n", " \"objective_target\": objective_target,\n", " \"scenario_techniques\": [ApiKeyTechnique.GetKey],\n", - " \"dataset_config\": ApiKeyDatasetConfiguration(dataset_names=ApiKey.required_datasets(), max_dataset_size=2),\n", + " \"dataset_config\": ApiKeyDatasetConfiguration(\n", + " sources=[DatasetSource(name=name) for name in ApiKey.required_datasets()], max_total=2\n", + " ),\n", " }\n", ")\n", "await api_key_scenario.initialize_async() # type: ignore\n", @@ -1159,7 +1161,7 @@ "actually asked for. A supplied `objective_scorer` replaces this fixed-trigger scorer; the\n", "harm family uses its separate `harm_scorer`. Caller technique converters run after the separators.\n", "\n", - "`max_dataset_size` is one budget before technique expansion. The default is 92 original\n", + "`max_total` is one budget before technique expansion. The default is 92 original\n", "groups, shared by six default techniques (552 execution units). Sampling reserves one group\n", "per selected family/trigger pair, then fills the remaining budget without replacement.\n", "A smaller budget than the number of pairs raises an error. An explicit dataset configuration\n", @@ -1228,7 +1230,9 @@ " \"objective_target\": objective_target,\n", " \"scenario_techniques\": [LatentInjectionTechnique.Bare],\n", " \"dataset_config\": LatentInjectionDatasetConfiguration(\n", - " dataset_names=LatentInjection.required_datasets(), families=[\"whois\"], max_dataset_size=1\n", + " sources=[DatasetSource(name=name) for name in LatentInjection.required_datasets()],\n", + " families=[\"whois\"],\n", + " max_total=1,\n", " ),\n", " }\n", ")\n", @@ -1357,7 +1361,7 @@ } ], "source": [ - "doctor_dataset_config = DatasetAttackConfiguration(dataset_names=[\"garak_doctor\"], max_dataset_size=1)\n", + "doctor_dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name=\"garak_doctor\")], max_total=1)\n", "\n", "doctor_scenario = Doctor()\n", "doctor_scenario.set_params_from_args( # type: ignore\n", @@ -1860,7 +1864,7 @@ ], "source": [ "audio_dataset_config = AudioAchillesHeelDatasetConfiguration(\n", - " dataset_names=[\"garak_audio_achilles_heel\"], max_dataset_size=1\n", + " sources=[DatasetSource(name=\"garak_audio_achilles_heel\")], max_total=1\n", ")\n", "\n", "audio_target = TargetRegistry.get_registry_singleton().instances.get(\"azure_openai_realtime\")\n", @@ -1988,7 +1992,7 @@ "\n", "**Available techniques:** `Repeat`, `DEFAULT`, and `ALL` all select the same probe.\n", "The default budget is 10 prompts across the entire dataset, not per word. Use\n", - "`DivergenceDatasetConfiguration(max_per_dataset=\"all\", max_total=\"all\", dataset_names=[\"garak_divergence\"])`\n", + "`DivergenceDatasetConfiguration(sources=[DatasetSource(name=\"garak_divergence\")], max_per_dataset=\"all\", max_total=\"all\")`\n", "to run all 36 prompts. The example below samples only two." ] }, @@ -2027,7 +2031,7 @@ " args={\n", " \"objective_target\": objective_target,\n", " \"scenario_techniques\": [DivergenceTechnique.Repeat],\n", - " \"dataset_config\": DivergenceDatasetConfiguration(dataset_names=[\"garak_divergence\"], max_dataset_size=2),\n", + " \"dataset_config\": DivergenceDatasetConfiguration(sources=[DatasetSource(name=\"garak_divergence\")], max_total=2),\n", " }\n", ")\n", "await divergence_scenario.initialize_async() # type: ignore\n", diff --git a/doc/scanner/garak.py b/doc/scanner/garak.py index 1fdbc0857b..6dc71866da 100644 --- a/doc/scanner/garak.py +++ b/doc/scanner/garak.py @@ -58,7 +58,7 @@ from pyrit.output import output_scenario_async from pyrit.registry import TargetRegistry -from pyrit.scenario import DatasetAttackConfiguration +from pyrit.scenario import DatasetAttackConfiguration, DatasetSource from pyrit.scenario.garak import ( ApiKey, ApiKeyDatasetConfiguration, @@ -103,8 +103,8 @@ "scenario_techniques": [PromptInjectTechnique.IgnorePrint], "goal_texts": ["I hate humans"], "dataset_config": PromptInjectDatasetConfiguration( - dataset_names=PromptInject.required_datasets(), - max_dataset_size=1, + sources=[DatasetSource(name=name) for name in PromptInject.required_datasets()], + max_total=1, ), } ) @@ -144,7 +144,7 @@ # > independently. # %% -dataset_config = EncodingDatasetConfiguration(dataset_names=["garak_slur_terms_en"], max_dataset_size=1) +dataset_config = EncodingDatasetConfiguration(sources=[DatasetSource(name="garak_slur_terms_en")], max_total=1) scenario = Encoding() scenario.set_params_from_args( # type: ignore @@ -187,7 +187,7 @@ # > `--include-baseline False` to omit this comparison. # %% -figstep_dataset_config = DatasetAttackConfiguration(dataset_names=["figstep"], max_dataset_size=1) +figstep_dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="figstep")], max_total=1) figstep_scenario = FigStep() figstep_scenario.set_params_from_args( # type: ignore @@ -315,7 +315,7 @@ # ``` # # **Available techniques:** `GetKey` and `CompleteKey`. `DEFAULT` and `ALL` both select the two -# techniques. `max_dataset_size` samples across all selected technique populations, not per service. +# techniques. `max_total` samples across all selected technique populations, not per service. # The base scenario persists the sample for resume. Use `ApiKeyDatasetConfiguration` with # `max_total="all"` to run all 348 requests. Standard technique converter stacks are supported. @@ -325,7 +325,9 @@ args={ "objective_target": objective_target, "scenario_techniques": [ApiKeyTechnique.GetKey], - "dataset_config": ApiKeyDatasetConfiguration(dataset_names=ApiKey.required_datasets(), max_dataset_size=2), + "dataset_config": ApiKeyDatasetConfiguration( + sources=[DatasetSource(name=name) for name in ApiKey.required_datasets()], max_total=2 + ), } ) await api_key_scenario.initialize_async() # type: ignore @@ -379,7 +381,7 @@ # actually asked for. A supplied `objective_scorer` replaces this fixed-trigger scorer; the # harm family uses its separate `harm_scorer`. Caller technique converters run after the separators. # -# `max_dataset_size` is one budget before technique expansion. The default is 92 original +# `max_total` is one budget before technique expansion. The default is 92 original # groups, shared by six default techniques (552 execution units). Sampling reserves one group # per selected family/trigger pair, then fills the remaining budget without replacement. # A smaller budget than the number of pairs raises an error. An explicit dataset configuration @@ -398,7 +400,9 @@ "objective_target": objective_target, "scenario_techniques": [LatentInjectionTechnique.Bare], "dataset_config": LatentInjectionDatasetConfiguration( - dataset_names=LatentInjection.required_datasets(), families=["whois"], max_dataset_size=1 + sources=[DatasetSource(name=name) for name in LatentInjection.required_datasets()], + families=["whois"], + max_total=1, ), } ) @@ -429,7 +433,7 @@ # tagged `default`, so `DEFAULT` and `ALL` currently coincide. # %% -doctor_dataset_config = DatasetAttackConfiguration(dataset_names=["garak_doctor"], max_dataset_size=1) +doctor_dataset_config = DatasetAttackConfiguration(sources=[DatasetSource(name="garak_doctor")], max_total=1) doctor_scenario = Doctor() doctor_scenario.set_params_from_args( # type: ignore @@ -560,7 +564,7 @@ # %% audio_dataset_config = AudioAchillesHeelDatasetConfiguration( - dataset_names=["garak_audio_achilles_heel"], max_dataset_size=1 + sources=[DatasetSource(name="garak_audio_achilles_heel")], max_total=1 ) audio_target = TargetRegistry.get_registry_singleton().instances.get("azure_openai_realtime") @@ -599,7 +603,7 @@ # # **Available techniques:** `Repeat`, `DEFAULT`, and `ALL` all select the same probe. # The default budget is 10 prompts across the entire dataset, not per word. Use -# `DivergenceDatasetConfiguration(max_per_dataset="all", max_total="all", dataset_names=["garak_divergence"])` +# `DivergenceDatasetConfiguration(sources=[DatasetSource(name="garak_divergence")], max_per_dataset="all", max_total="all")` # to run all 36 prompts. The example below samples only two. # %% @@ -608,7 +612,7 @@ args={ "objective_target": objective_target, "scenario_techniques": [DivergenceTechnique.Repeat], - "dataset_config": DivergenceDatasetConfiguration(dataset_names=["garak_divergence"], max_dataset_size=2), + "dataset_config": DivergenceDatasetConfiguration(sources=[DatasetSource(name="garak_divergence")], max_total=2), } ) await divergence_scenario.initialize_async() # type: ignore diff --git a/pyrit/scenario/scenarios/airt/psychosocial.py b/pyrit/scenario/scenarios/airt/psychosocial.py index be50a0d2f4..9bc365314f 100644 --- a/pyrit/scenario/scenarios/airt/psychosocial.py +++ b/pyrit/scenario/scenarios/airt/psychosocial.py @@ -494,7 +494,8 @@ def _validate_runtime_configuration(self) -> None: if cap != "all" and cap < len(self._selected_sub_harms()): raise DatasetConstraintError( f"Psychosocial max_total ({cap}) must cover every selected sub-harm " - f"({len(self._selected_sub_harms())}); use max_per_dataset=1 for one objective per sub-harm." + f"({len(self._selected_sub_harms())}). Increase the total limit to at least " + f"{len(self._selected_sub_harms())}, or select one sub-harm." ) async def _resolve_seed_groups_by_dataset_async( diff --git a/tests/unit/scenario/airt/test_psychosocial.py b/tests/unit/scenario/airt/test_psychosocial.py index 79820acff0..816fb42e0d 100644 --- a/tests/unit/scenario/airt/test_psychosocial.py +++ b/tests/unit/scenario/airt/test_psychosocial.py @@ -260,7 +260,10 @@ async def test_backend_total_one_rejects_both_sub_harms_before_reads_async( with ( patch.object(DatasetAttackConfiguration, "prepare_async", side_effect=AssertionError("No preparation")), patch.object(CentralMemory.get_memory_instance(), "get_seeds_async", side_effect=AssertionError("No reads")), - pytest.raises(DatasetConstraintError, match="max_total.*cover every selected sub-harm"), + pytest.raises( + DatasetConstraintError, + match=r"max_total \(1\).*Increase the total limit to at least 2, or select one sub-harm\.", + ), ): if preview: await scenario.get_run_size_estimate_async()