Skip to content

Commit 484f0d3

Browse files
Add input/output modality validation to TargetRequirements.validate() (#1778)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent c94d334 commit 484f0d3

5 files changed

Lines changed: 169 additions & 2 deletions

File tree

doc/code/targets/0_prompt_targets.md

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,19 @@ CHAT_TARGET_REQUIREMENTS.validate(target=target)
7878

7979
`TargetRequirements.validate` collects every missing capability and raises a single `ValueError`. For one-off checks against a single capability you can also call `target.configuration.ensure_can_handle(capability=...)` directly.
8080

81+
`TargetRequirements` can also enforce **modality** constraints via `required_input_modalities` and `required_output_modalities`. Each entry is a set of `PromptDataType` values the consumer needs the target to accept (or produce). At least one of the target's modality combos must be a superset of each required combo:
82+
83+
```python
84+
from pyrit.prompt_target import TargetRequirements
85+
86+
# A consumer that requires image input and text output
87+
VISION_REQUIREMENTS = TargetRequirements(
88+
required_input_modalities=frozenset({frozenset({"image_path"})}),
89+
required_output_modalities=frozenset({frozenset({"text"})}),
90+
)
91+
VISION_REQUIREMENTS.validate(target=target)
92+
```
93+
8194
### Adapting vs raising
8295

8396
Some capability gaps can be papered over by PyRIT itself. For example, a single-turn target can be made to *appear* multi-turn by flattening the conversation history into a single prompt before sending. The `CapabilityHandlingPolicy` controls this on a per-capability basis:

doc/code/targets/6_1_target_capabilities.ipynb

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -154,7 +154,12 @@
154154
"multi-turn + editable history — the replacement for the deprecated `PromptChatTarget` type check.\n",
155155
"\n",
156156
"`TargetRequirements.validate` collects every missing capability and raises a single `ValueError` so\n",
157-
"callers see all violations at once."
157+
"callers see all violations at once.\n",
158+
"\n",
159+
"`TargetRequirements` can also enforce **modality** constraints via `required_input_modalities` and\n",
160+
"`required_output_modalities`. Each entry is a set of `PromptDataType` values the consumer needs\n",
161+
"the target to accept (or produce). At least one of the target's modality combos must be a superset\n",
162+
"of each required combo."
158163
]
159164
},
160165
{

doc/code/targets/6_1_target_capabilities.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -92,6 +92,11 @@
9292
#
9393
# `TargetRequirements.validate` collects every missing capability and raises a single `ValueError` so
9494
# callers see all violations at once.
95+
#
96+
# `TargetRequirements` can also enforce **modality** constraints via `required_input_modalities` and
97+
# `required_output_modalities`. Each entry is a set of `PromptDataType` values the consumer needs
98+
# the target to accept (or produce). At least one of the target's modality combos must be a superset
99+
# of each required combo.
95100

96101
# %%
97102
from pyrit.prompt_target import CHAT_TARGET_REQUIREMENTS

pyrit/prompt_target/common/target_requirements.py

Lines changed: 44 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from pyrit.prompt_target.common.target_capabilities import CapabilityName
1010

1111
if TYPE_CHECKING:
12+
from pyrit.models import PromptDataType
1213
from pyrit.prompt_target.common.prompt_target import PromptTarget
1314

1415

@@ -33,10 +34,20 @@ class TargetRequirements:
3334
consumer's semantics (e.g. an attack that depends on the target
3435
remembering prior turns, where history-squash normalization would
3536
collapse the conversation into a single prompt).
37+
38+
Modality requirements are also supported:
39+
40+
* ``required_input_modalities`` — each entry is a frozenset of
41+
:class:`PromptDataType` values the consumer needs the target to
42+
accept. At least one of the target's input modality combos must be
43+
a superset of each required combo.
44+
* ``required_output_modalities`` — same semantics for outputs.
3645
"""
3746

3847
required: frozenset[CapabilityName] = field(default_factory=frozenset)
3948
native_required: frozenset[CapabilityName] = field(default_factory=frozenset)
49+
required_input_modalities: frozenset[frozenset[PromptDataType]] = field(default_factory=frozenset)
50+
required_output_modalities: frozenset[frozenset[PromptDataType]] = field(default_factory=frozenset)
4051

4152
def validate(self, *, target: PromptTarget) -> None:
4253
"""
@@ -52,7 +63,9 @@ def validate(self, *, target: PromptTarget) -> None:
5263
Raises:
5364
ValueError: If any ``native_required`` capability is not natively
5465
supported, or if any ``required`` capability is not supported
55-
natively and has no ``ADAPT`` entry in the target's policy.
66+
natively and has no ``ADAPT`` entry in the target's policy,
67+
or if the target's modalities do not satisfy
68+
``required_input_modalities`` / ``required_output_modalities``.
5669
"""
5770
errors: list[str] = [
5871
f"Target must natively support '{capability.value}'; adaptation is not acceptable for this consumer."
@@ -66,12 +79,42 @@ def validate(self, *, target: PromptTarget) -> None:
6679
except ValueError as exc:
6780
errors.append(str(exc))
6881

82+
errors.extend(
83+
self._check_modalities(
84+
required=self.required_input_modalities,
85+
supported=target.configuration.capabilities.input_modalities,
86+
direction="input",
87+
)
88+
)
89+
errors.extend(
90+
self._check_modalities(
91+
required=self.required_output_modalities,
92+
supported=target.configuration.capabilities.output_modalities,
93+
direction="output",
94+
)
95+
)
96+
6997
if errors:
7098
raise ValueError(
7199
f"Target does not satisfy {len(errors)} required capability(ies):\n"
72100
+ "\n".join(f" - {e}" for e in errors)
73101
)
74102

103+
@staticmethod
104+
def _check_modalities(
105+
*,
106+
required: frozenset[frozenset[PromptDataType]],
107+
supported: frozenset[frozenset[PromptDataType]],
108+
direction: str,
109+
) -> list[str]:
110+
"""Return error strings for each required modality combo not covered by *supported*."""
111+
return [
112+
f"Target must support {direction} modality {{{', '.join(sorted(combo))}}}; "
113+
f"supported: {[sorted(s) for s in sorted(supported, key=lambda s: sorted(s))]}."
114+
for combo in sorted(required, key=lambda c: sorted(c))
115+
if not any(combo <= sup for sup in supported)
116+
]
117+
75118

76119
def _build_chat_target_requirements() -> TargetRequirements:
77120
"""

tests/unit/prompt_target/target/test_target_requirements.py

Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -262,3 +262,104 @@ def test_validate_aggregates_all_violations():
262262
assert "2 required capability" in message
263263
assert CapabilityName.MULTI_TURN.value in message
264264
assert CapabilityName.JSON_OUTPUT.value in message
265+
266+
267+
# ---------------------------------------------------------------------------
268+
# Modality validation
269+
# ---------------------------------------------------------------------------
270+
271+
272+
def test_validate_passes_when_no_modality_requirements():
273+
"""Backward compat: empty modality requirements should always pass."""
274+
target = _make_target(
275+
configuration=TargetConfiguration(
276+
capabilities=TargetCapabilities(),
277+
),
278+
)
279+
TargetRequirements().validate(target=target)
280+
281+
282+
def test_validate_passes_when_input_modality_matches():
283+
reqs = TargetRequirements(
284+
required_input_modalities=frozenset({frozenset({"text"})}),
285+
)
286+
target = _make_target(
287+
configuration=TargetConfiguration(
288+
capabilities=TargetCapabilities(
289+
input_modalities=frozenset({frozenset({"text"})}),
290+
),
291+
),
292+
)
293+
reqs.validate(target=target)
294+
295+
296+
def test_validate_passes_when_target_modality_is_superset():
297+
reqs = TargetRequirements(
298+
required_input_modalities=frozenset({frozenset({"text"})}),
299+
)
300+
target = _make_target(
301+
configuration=TargetConfiguration(
302+
capabilities=TargetCapabilities(
303+
input_modalities=frozenset({frozenset({"text", "image_path"})}),
304+
),
305+
),
306+
)
307+
reqs.validate(target=target)
308+
309+
310+
def test_validate_fails_on_missing_input_modality():
311+
reqs = TargetRequirements(
312+
required_input_modalities=frozenset({frozenset({"image_path"})}),
313+
)
314+
target = _make_target(
315+
configuration=TargetConfiguration(
316+
capabilities=TargetCapabilities(
317+
input_modalities=frozenset({frozenset({"text"})}),
318+
),
319+
),
320+
)
321+
with pytest.raises(ValueError, match="input modality"):
322+
reqs.validate(target=target)
323+
324+
325+
def test_validate_fails_on_missing_output_modality():
326+
reqs = TargetRequirements(
327+
required_output_modalities=frozenset({frozenset({"audio_path"})}),
328+
)
329+
target = _make_target(
330+
configuration=TargetConfiguration(
331+
capabilities=TargetCapabilities(
332+
output_modalities=frozenset({frozenset({"text"})}),
333+
),
334+
),
335+
)
336+
with pytest.raises(ValueError, match="output modality"):
337+
reqs.validate(target=target)
338+
339+
340+
def test_validate_aggregates_modality_and_capability_errors():
341+
reqs = TargetRequirements(
342+
native_required=frozenset({CapabilityName.MULTI_TURN}),
343+
required_input_modalities=frozenset({frozenset({"image_path"})}),
344+
)
345+
target = _make_target(
346+
configuration=TargetConfiguration(
347+
capabilities=TargetCapabilities(
348+
supports_multi_turn=False,
349+
input_modalities=frozenset({frozenset({"text"})}),
350+
),
351+
),
352+
)
353+
with pytest.raises(ValueError) as exc_info:
354+
reqs.validate(target=target)
355+
356+
message = str(exc_info.value)
357+
assert "2 required capability" in message
358+
assert CapabilityName.MULTI_TURN.value in message
359+
assert "input modality" in message
360+
361+
362+
def test_default_modality_requirements_are_empty():
363+
reqs = TargetRequirements()
364+
assert reqs.required_input_modalities == frozenset()
365+
assert reqs.required_output_modalities == frozenset()

0 commit comments

Comments
 (0)