diff --git a/pyrit/datasets/seed_datasets/seed_dataset_provider.py b/pyrit/datasets/seed_datasets/seed_dataset_provider.py index 45fb042ee9..be58c21ab5 100644 --- a/pyrit/datasets/seed_datasets/seed_dataset_provider.py +++ b/pyrit/datasets/seed_datasets/seed_dataset_provider.py @@ -222,9 +222,18 @@ def _match_single_criterion( filter_vals = getattr(criterion, field.name) meta_vals = getattr(metadata, field.name) - if filter_vals is None or meta_vals is None: + if filter_vals is None: continue + # `meta_vals is None` means the dataset never declared this axis, which + # is not the same as declaring it and matching. Skipping here dropped + # the filtered axis entirely, so every dataset silent on that axis came + # back as a match. It also disagreed with get_all_dataset_names_async, + # which excludes datasets carrying no metadata at all - a dataset with + # partial metadata was treated as better qualified than one with none. + if meta_vals is None: + return False + if strict_match: if filter_vals - meta_vals: return False diff --git a/tests/unit/datasets/test_seed_dataset_provider.py b/tests/unit/datasets/test_seed_dataset_provider.py index 12b3c3bec6..3110b95528 100644 --- a/tests/unit/datasets/test_seed_dataset_provider.py +++ b/tests/unit/datasets/test_seed_dataset_provider.py @@ -371,6 +371,39 @@ def test_modalities(self): dataset_filter=SeedDatasetFilter(modalities={"audio"}), ) + def test_undeclared_axis_does_not_match(self): + """A dataset silent on a filtered axis does not satisfy that axis.""" + metadata = SeedDatasetMetadata(tags={"safety"}, size={"small"}, source_type={"local"}) + assert not SeedDatasetProvider._match_filter_to_metadata( + metadata=metadata, + dataset_filter=SeedDatasetFilter(modalities={"audio"}), + ) + assert not SeedDatasetProvider._match_filter_to_metadata( + metadata=metadata, + dataset_filter=SeedDatasetFilter(harm_categories={"violence"}), + ) + # An axis the dataset does declare is still matched normally. + assert SeedDatasetProvider._match_filter_to_metadata( + metadata=metadata, + dataset_filter=SeedDatasetFilter(tags={"safety"}), + ) + + def test_undeclared_axis_does_not_match_strict(self): + """strict_match also rejects a dataset silent on the filtered axis.""" + metadata = SeedDatasetMetadata(tags={"safety"}, size={"small"}, source_type={"local"}) + assert not SeedDatasetProvider._match_filter_to_metadata( + metadata=metadata, + dataset_filter=SeedDatasetFilter(modalities={"audio"}, strict_match=True), + ) + + def test_undeclared_axis_does_not_mask_a_declared_mismatch(self): + """Two filtered axes, one declared and mismatching: still no match.""" + metadata = SeedDatasetMetadata(modalities={"text"}) + assert not SeedDatasetProvider._match_filter_to_metadata( + metadata=metadata, + dataset_filter=SeedDatasetFilter(modalities={"text"}, harm_categories={"violence"}), + ) + def test_sources(self): """Source filter checks membership.""" metadata = SeedDatasetMetadata(source_type={"remote"})