diff --git a/doc/scanner/garak.ipynb b/doc/scanner/garak.ipynb index 0425b4e870..18cc5a8c90 100644 --- a/doc/scanner/garak.ipynb +++ b/doc/scanner/garak.ipynb @@ -620,7 +620,7 @@ "**CLI example:**\n", "\n", "```bash\n", - "pyrit_scan run garak.web_injection --target openai_chat --techniques xss --max-dataset-size 1\n", + "pyrit_scan run garak.web_injection --target openai_chat --techniques xss\n", "```\n", "\n", "**Available techniques** (8 probes): MarkdownImageExfil, ColabAIDataLeakage,\n", diff --git a/doc/scanner/garak.py b/doc/scanner/garak.py index 38673b15cd..23c41e5555 100644 --- a/doc/scanner/garak.py +++ b/doc/scanner/garak.py @@ -221,7 +221,7 @@ # **CLI example:** # # ```bash -# pyrit_scan run garak.web_injection --target openai_chat --techniques xss --max-dataset-size 1 +# pyrit_scan run garak.web_injection --target openai_chat --techniques xss # ``` # # **Available techniques** (8 probes): MarkdownImageExfil, ColabAIDataLeakage, diff --git a/pyrit/backend/services/scenario_configuration_resolver.py b/pyrit/backend/services/scenario_configuration_resolver.py index b0c6e19256..361b8c203e 100644 --- a/pyrit/backend/services/scenario_configuration_resolver.py +++ b/pyrit/backend/services/scenario_configuration_resolver.py @@ -148,6 +148,11 @@ def resolve_configuration( resolved["technique_converters"] = technique_converters if dataset_names or has_total_override or filters: + cls._validate_dataset_overrides( + scenario_name=scenario_name, + scenario=introspection_instance, + has_total_override=has_total_override, + ) default_config = introspection_instance._default_dataset_config config = default_config.with_overrides(filters=filters) if dataset_names: @@ -163,6 +168,30 @@ def resolve_configuration( return resolved + @staticmethod + def _validate_dataset_overrides(*, scenario_name: str, scenario: Scenario, has_total_override: bool) -> None: + """ + Reject dataset overrides the scenario would not apply. + + Args: + scenario_name: Scenario name used in error messages. + scenario: Introspection instance of the scenario. + has_total_override: Whether a non-default total limit was requested. + + Raises: + ValueError: If the scenario has a fixed dataset, or ignores the total limit. + """ + if not any(parameter.name == "dataset_config" for parameter in scenario.supported_parameters()): + raise ValueError( + f"Scenario '{scenario_name}' uses a fixed dataset, so it doesn't take dataset names, " + "dataset filters or a dataset size limit." + ) + if has_total_override and not scenario.USES_DATASET_SIZE_LIMIT: + raise ValueError( + f"Scenario '{scenario_name}' doesn't use a dataset size limit; its own settings decide " + "how many prompts it runs. Leave the dataset size at its default." + ) + @classmethod def resolve_techniques_and_converters( cls, diff --git a/tests/unit/backend/test_scenario_configuration_resolver.py b/tests/unit/backend/test_scenario_configuration_resolver.py index 22ff1e2c0d..3d1c9ba8f6 100644 --- a/tests/unit/backend/test_scenario_configuration_resolver.py +++ b/tests/unit/backend/test_scenario_configuration_resolver.py @@ -3,7 +3,7 @@ """Adversarial target resolution validates without changing execution scopes.""" -from typing import Literal +from typing import Any, Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -23,7 +23,11 @@ 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.exploitation import Exploitation +from pyrit.scenario.scenarios.garak.package_hallucination import PackageHallucination from pyrit.scenario.scenarios.garak.prompt_inject import PromptInject, PromptInjectDatasetConfiguration +from pyrit.scenario.scenarios.garak.system_prompt_extraction import SystemPromptExtraction +from pyrit.scenario.scenarios.garak.web_injection import WebInjection from pyrit.score import TrueFalseScorer from unit.mocks import MockPromptTarget @@ -183,3 +187,45 @@ def test_resolve_adversarial_target_rejects_invalid_selection_without_changing_s with pytest.raises(ValueError, match=message): ScenarioConfigurationResolver.resolve_adversarial_target(target_name=selection) assert get_default_adversarial_target() is outer + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("scenario_class", [SystemPromptExtraction, WebInjection, PackageHallucination]) +@pytest.mark.parametrize("limit", [3, "all"]) +def test_total_limit_is_rejected_when_the_scenario_ignores_it( + *, scenario_class: type[Scenario], limit: DatasetLimit +) -> None: + with ( + patch.object(Scenario, "_get_default_objective_scorer", return_value=MagicMock(spec=TrueFalseScorer)), + pytest.raises(ValueError, match="doesn't use a dataset size limit"), + ): + ScenarioConfigurationResolver.resolve_configuration( + scenario_name="test", scenario_class=scenario_class, max_dataset_size=limit + ) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("scenario_class", [SystemPromptExtraction, WebInjection, PackageHallucination]) +def test_dataset_names_still_reach_a_scenario_that_ignores_the_total_limit(scenario_class: type[Scenario]) -> None: + with patch.object(Scenario, "_get_default_objective_scorer", return_value=MagicMock(spec=TrueFalseScorer)): + names = scenario_class()._default_dataset_config.dataset_names + resolved = ScenarioConfigurationResolver.resolve_configuration( + scenario_name="test", scenario_class=scenario_class, dataset_names=names, max_dataset_size="default" + ) + assert resolved["dataset_config"].dataset_names == names + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "overrides", + [ + {"max_dataset_size": 3}, + {"dataset_names": ["garak_exploitation_sql_injection"]}, + {"dataset_filters": {"data_types": ["text"]}}, + ], +) +def test_dataset_overrides_are_rejected_for_a_fixed_dataset_scenario(overrides: dict[str, Any]) -> None: + with pytest.raises(ValueError, match="uses a fixed dataset"): + ScenarioConfigurationResolver.resolve_configuration( + scenario_name="garak.exploitation", scenario_class=Exploitation, **overrides + ) diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index dd17a8661c..8eaf6aec00 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -301,6 +301,7 @@ def mock_all_registries(mock_memory): mock_scenario_class = MagicMock(return_value=mock_scenario_instance) mock_scenario_instance._technique_class = MagicMock() mock_scenario_instance._default_dataset_config = MagicMock() + mock_scenario_instance.supported_parameters.return_value = Scenario.supported_parameters() mock_sr = MagicMock() mock_sr.get_class.return_value = mock_scenario_class diff --git a/tests/unit/backend/test_scenario_service.py b/tests/unit/backend/test_scenario_service.py index 19dbdb2883..578f840fc2 100644 --- a/tests/unit/backend/test_scenario_service.py +++ b/tests/unit/backend/test_scenario_service.py @@ -364,6 +364,7 @@ def construct() -> MagicMock: introspected.append(get_default_adversarial_target()) scenario = MagicMock(spec=Scenario) scenario._default_dataset_config = DatasetAttackConfiguration(dataset_names=["test_dataset"]) + scenario.supported_parameters.return_value = Scenario.supported_parameters() return scenario async def estimate_async(**kwargs: object) -> ScenarioRunSizeEstimate: @@ -1543,6 +1544,7 @@ async def test_estimate_resolves_total_limit_async( original = DatasetAttackConfiguration(dataset_names=["harmbench"], max_total=20) scenario_class = MagicMock() scenario_class.return_value._default_dataset_config = original + scenario_class.return_value.supported_parameters.return_value = Scenario.supported_parameters() with patch.object(ScenarioService, "__init__", lambda self: None): service = ScenarioService() service._registry = MagicMock() @@ -1584,6 +1586,7 @@ async def test_configured_estimate_uses_shared_launch_resolution(self) -> None: introspection_instance = MagicMock() introspection_instance._technique_class = _EstimateTechnique introspection_instance._default_dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"]) + introspection_instance.supported_parameters.return_value = Scenario.supported_parameters() scenario_class = MagicMock(return_value=introspection_instance) scenario_class.supported_parameters.return_value = [ Parameter(name=name, description="", param_type=int) @@ -1643,6 +1646,7 @@ async def test_configured_estimate_rejects_incompatible_v4_jailbreak_technique(s introspection_instance = MagicMock() introspection_instance._technique_class = _EstimateTechnique introspection_instance._default_dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"]) + introspection_instance.supported_parameters.return_value = Scenario.supported_parameters() scenario_class = MagicMock(return_value=introspection_instance) with patch.object(ScenarioService, "__init__", lambda self: None): @@ -1669,6 +1673,7 @@ async def test_configured_estimate_without_target_does_not_resolve_or_send_to_ta introspection_instance = MagicMock() introspection_instance._technique_class = _EstimateTechnique introspection_instance._default_dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"]) + introspection_instance.supported_parameters.return_value = Scenario.supported_parameters() scenario_class = MagicMock(return_value=introspection_instance) with (