Skip to content
Open
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions pyrit/datasets/seed_datasets/seed_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,7 +184,7 @@ class SeedDatasetFilter:
SeedDatasetMetadata(size={"large"}, modalities={"image"}),
])

Passing both flat kwargs and criteria raises ValueError.
Passing both flat kwargs and criteria, or an empty criteria list, raises ValueError.

Special tags:
- "all": Returns every dataset, ignores all other fields. This tag will
Expand All @@ -193,7 +193,7 @@ class SeedDatasetFilter:
strict_match=True, loses its shortcut and is treated as a normal tag.

Args:
criteria: Explicit list of SeedDatasetMetadata to OR-match against.
criteria: Non-empty list of SeedDatasetMetadata to OR-match against.
strict_match: If True, within-axis matching uses AND (all filter values
must be present) instead of OR (any overlap suffices).
**kwargs: Flat metadata fields (tags, size, modalities, etc.) for simple use.
Expand Down Expand Up @@ -221,12 +221,12 @@ def __init__(
])

Args:
criteria: Explicit list of SeedDatasetMetadata to OR-match against.
criteria: Non-empty list of SeedDatasetMetadata to OR-match against.
strict_match: If True, within-axis matching uses AND instead of OR.
**kwargs: Flat metadata fields passed to SeedDatasetMetadata.

Raises:
ValueError: If both criteria and flat kwargs are provided.
ValueError: If both criteria and flat kwargs are provided, or criteria is empty.
"""
if criteria is not None and kwargs:
raise ValueError("Cannot pass both 'criteria' and flat metadata kwargs. Use one or the other.")
Expand All @@ -248,6 +248,11 @@ def _normalize_criterion(c: SeedDatasetMetadata) -> SeedDatasetMetadata:

self.criteria = [_normalize_criterion(c) for c in self.criteria]

if not self.criteria:
raise ValueError(
"'criteria' must contain at least one criterion. Omit 'criteria' for an unconstrained filter."
)

self.strict_match = strict_match
self._validate()

Expand Down
15 changes: 15 additions & 0 deletions tests/unit/datasets/test_seed_dataset_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -371,3 +371,18 @@ def test_tags_values(self):
def test_harm_categories_values(self):
f = SeedDatasetFilter(harm_categories={"violence", "cybercrime"})
assert "violence" in f.criteria[0].harm_categories


@pytest.mark.parametrize("strict_match", [False, True])
def test_empty_criteria_raises(strict_match: bool) -> None:
with pytest.raises(ValueError, match="at least one criterion"):
SeedDatasetFilter(criteria=[], strict_match=strict_match)


def test_empty_criteria_with_kwargs_preserves_conflicting_arguments_error() -> None:
with pytest.raises(ValueError, match="Cannot pass both"):
SeedDatasetFilter(criteria=[], tags={"safety"})


def test_unconstrained_criterion_preserves_default_filter() -> None:
assert SeedDatasetFilter(criteria=[SeedDatasetMetadata()]).criteria == SeedDatasetFilter().criteria
Loading