Skip to content

Commit b7af0de

Browse files
author
Victor Valbuena
committed
FEAT: Add TargetRegistry.get_by_tag_query for TagQuery-based lookup
1 parent 42887b9 commit b7af0de

2 files changed

Lines changed: 90 additions & 0 deletions

File tree

pyrit/registry/object_registries/target_registry.py

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,9 +12,11 @@
1212
import logging
1313
from typing import TYPE_CHECKING, Optional, Union
1414

15+
from pyrit.registry.object_registries.base_instance_registry import RegistryEntry
1516
from pyrit.registry.object_registries.retrievable_instance_registry import (
1617
RetrievableInstanceRegistry,
1718
)
19+
from pyrit.registry.tag_query import TagQuery
1820

1921
if TYPE_CHECKING:
2022
from pyrit.prompt_target import PromptTarget
@@ -74,3 +76,30 @@ def get_instance_by_name(self, name: str) -> Optional[PromptTarget]:
7476
The target instance, or None if not found.
7577
"""
7678
return self.get(name)
79+
80+
def get_by_tag_query(self, *, query: TagQuery) -> list[RegistryEntry[PromptTarget]]:
81+
"""
82+
Get all entries whose tag keys satisfy ``query``.
83+
84+
``TagQuery`` operates on a tag set, so this method matches against
85+
``entry.tags.keys()`` and ignores tag values. For value-aware
86+
single-tag lookups use ``get_by_tag(*, tag, value)`` on the base
87+
class.
88+
89+
Composite queries compose with ``&`` and ``|`` operators, e.g.
90+
``TagQuery.all("adversarial") & TagQuery.any_of("singleturn", "multiturn")``.
91+
92+
Args:
93+
query: The tag predicate to evaluate against each entry.
94+
95+
Returns:
96+
List of matching ``RegistryEntry`` objects sorted by registry name.
97+
"""
98+
results: list[RegistryEntry[PromptTarget]] = []
99+
# Note: this erases insertion order, but respects the base_instance_registry pattern
100+
# (get_by_tag).
101+
for name in sorted(self._registry_items.keys()):
102+
entry = self._registry_items[name]
103+
if query.matches(set(entry.tags.keys())):
104+
results.append(entry)
105+
return results

tests/unit/registry/test_target_registry.py

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -230,3 +230,64 @@ def test_list_metadata_filter_by_class_name(self):
230230
assert len(mock_metadata) == 2
231231
for m in mock_metadata:
232232
assert m.class_name == "MockPromptTarget"
233+
234+
235+
@pytest.mark.usefixtures("patch_central_database")
236+
class TestTargetRegistryGetByTagQuery:
237+
"""Tests for ``TargetRegistry.get_by_tag_query`` (TagQuery-aware tag lookup)."""
238+
239+
def setup_method(self):
240+
"""Reset and populate a fresh registry for each test."""
241+
TargetRegistry.reset_instance()
242+
self.registry = TargetRegistry.get_registry_singleton()
243+
244+
self.registry.register_instance(MockPromptTarget(), name="adv_single", tags=["adversarial", "singleturn"])
245+
self.registry.register_instance(MockPromptTarget(), name="adv_multi", tags=["adversarial", "multiturn"])
246+
self.registry.register_instance(MockPromptChatTarget(), name="scorer_only", tags=["scorer"])
247+
self.registry.register_instance(MockPromptTarget(), name="untagged")
248+
249+
def teardown_method(self):
250+
"""Reset the singleton after each test."""
251+
TargetRegistry.reset_instance()
252+
253+
def test_get_by_tag_query_returns_matching(self):
254+
"""A leaf ``TagQuery.all`` returns every entry whose tag set contains the required tag."""
255+
from pyrit.registry.tag_query import TagQuery
256+
257+
results = self.registry.get_by_tag_query(query=TagQuery.all("adversarial"))
258+
259+
names = [entry.name for entry in results]
260+
assert names == ["adv_multi", "adv_single"]
261+
262+
def test_get_by_tag_query_empty(self):
263+
"""A query that matches no entries returns an empty list (not raise)."""
264+
from pyrit.registry.tag_query import TagQuery
265+
266+
results = self.registry.get_by_tag_query(query=TagQuery.all("nonexistent_tag"))
267+
assert results == []
268+
269+
def test_get_by_tag_query_composite_and_or(self):
270+
"""Composite queries via ``&`` / ``|`` evaluate as expected."""
271+
from pyrit.registry.tag_query import TagQuery
272+
273+
query = TagQuery.all("adversarial") & TagQuery.any_of("singleturn", "multiturn")
274+
results = self.registry.get_by_tag_query(query=query)
275+
276+
names = [entry.name for entry in results]
277+
assert names == ["adv_multi", "adv_single"]
278+
279+
narrower = TagQuery.all("adversarial") & TagQuery.any_of("singleturn")
280+
narrow_names = [entry.name for entry in self.registry.get_by_tag_query(query=narrower)]
281+
assert narrow_names == ["adv_single"]
282+
283+
def test_get_by_tag_query_matches_keys_not_values(self):
284+
"""``TagQuery`` evaluates against tag keys; tag values are ignored by this method."""
285+
from pyrit.registry.tag_query import TagQuery
286+
287+
self.registry.add_tags(name="adv_single", tags={"priority": "high"})
288+
289+
priority_matches = self.registry.get_by_tag_query(query=TagQuery.all("priority"))
290+
assert [entry.name for entry in priority_matches] == ["adv_single"]
291+
292+
value_lookup = self.registry.get_by_tag_query(query=TagQuery.all("high"))
293+
assert value_lookup == []

0 commit comments

Comments
 (0)