Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
47 changes: 37 additions & 10 deletions pyrit/models/target/target_capabilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from enum import Enum
from typing import cast

from pydantic import BaseModel, ConfigDict, Field, computed_field
from pydantic import BaseModel, ConfigDict, Field, computed_field, field_serializer

from pyrit.models.literals import PromptDataType # noqa: TC001 (runtime-required by Pydantic field annotations)

Expand Down Expand Up @@ -65,10 +65,13 @@ class attribute. Users can override individual capabilities per instance
across targets and reused as a known-model profile.

This model also serves as the REST wire snapshot of a target's capabilities
(it is embedded in ``TargetInstance``). The modality *combination* fields
(``input_modalities`` / ``output_modalities``) are excluded from serialization;
API consumers read the flattened ``supported_input_modalities`` /
``supported_output_modalities`` computed fields instead.
(it is embedded in ``TargetInstance``), so serialization must be lossless: a
dumped payload validated back in has to rebuild an equal object. The immutable
``frozenset[frozenset]`` combination fields are therefore serialized in a
deterministic form -- a sorted list of sorted lists -- which pydantic parses back
into the same structure. The flattened ``supported_input_modalities`` /
``supported_output_modalities`` computed fields are emitted alongside them for
API consumers that only need per-piece modality checks.
"""

model_config = ConfigDict(frozen=True)
Expand Down Expand Up @@ -104,13 +107,35 @@ class attribute. Users can override individual capabilities per instance
supports_streaming_audio: bool = False

#: The input modalities supported by the target, as combinations of data types
#: (e.g., ``{{"text"}, {"image_path", "text"}}``). Excluded from serialization —
#: API consumers read the flattened ``supported_input_modalities`` instead.
input_modalities: frozenset[frozenset[PromptDataType]] = Field(default=_DEFAULT_TEXT_MODALITIES, exclude=True)
#: (e.g., ``{{"text"}, {"image_path", "text"}}``). Serialized as a sorted list of
#: sorted lists — see ``_sorted_modality_combinations``.
input_modalities: frozenset[frozenset[PromptDataType]] = Field(default=_DEFAULT_TEXT_MODALITIES)

#: The output modalities supported by the target, as combinations of data types.
#: Excluded from serialization — see ``supported_output_modalities``.
output_modalities: frozenset[frozenset[PromptDataType]] = Field(default=_DEFAULT_TEXT_MODALITIES, exclude=True)
#: Serialized as a sorted list of sorted lists — see ``_sorted_modality_combinations``.
output_modalities: frozenset[frozenset[PromptDataType]] = Field(default=_DEFAULT_TEXT_MODALITIES)

@field_serializer("input_modalities", "output_modalities")
def _sorted_modality_combinations(self, combinations: frozenset[frozenset[PromptDataType]]) -> list[list[str]]:
"""
Serialize the modality combinations into a deterministic wire form.

``frozenset`` has no order, so dumping it directly would emit an
arbitrary permutation and make the payload unstable across processes —
which matters because this model is the REST snapshot embedded in
``TargetInstance``, and because equal capabilities must produce equal
payloads. Sorting each combination's data types and then sorting the
combinations themselves removes that nondeterminism. Pydantic parses
the result back into the same ``frozenset[frozenset]`` structure, so
the round trip is lossless.

Args:
combinations: The modality combinations to serialize.

Returns:
list[list[str]]: The combinations as sorted lists of sorted data types.
"""
return sorted(sorted(str(data_type) for data_type in combination) for combination in combinations)

@computed_field( # type: ignore[prop-decorator]
description="Sorted unique input modality data types the target accepts (e.g., ['image_path', 'text'])",
Expand All @@ -123,6 +148,8 @@ def supported_input_modalities(self) -> list[str]:
The internal ``input_modalities`` models modality *combinations*
(``frozenset[frozenset]``); API consumers use only per-piece modality
checks, so this flattens the combinations into a sorted unique list.
This projection is derived, not stored — validating the wire form
rebuilds the combinations, and this property follows from them.

Returns:
list[str]: Sorted unique input modality data types.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2123,7 +2123,7 @@ async def stop_async() -> None:
await asyncio.wait_for(target.cleanup_target_async(), timeout=2.0)


@pytest.mark.usefixtures("patch_central_database")
@pytest.mark.usefixtures("patch_central_database", "mock_copilot_startup_io")
async def test_normalizer_stops_owned_client_when_startup_is_cancelled_async(
*,
sdk: Any,
Expand Down
170 changes: 170 additions & 0 deletions tests/unit/prompt_target/target/test_target_capabilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,10 @@
from unittest.mock import patch

import pytest
from pydantic import ValidationError

from pyrit.models.catalog import TargetInstance
from pyrit.models.identifiers import TargetIdentifier
from pyrit.prompt_target.common.conversation_normalization_pipeline import NORMALIZABLE_CAPABILITIES
from pyrit.prompt_target.common.target_capabilities import (
CapabilityHandlingPolicy,
Expand Down Expand Up @@ -561,3 +564,170 @@ def test_prompt_target_preserves_system_prompt_for_recognized_model(self):
assert result.capabilities.supports_multi_turn is True
assert result.capabilities.supports_multi_message_pieces is True
assert result.capabilities.supports_system_prompt is True


class TestTargetCapabilitiesWireRoundTrip:
"""Test that the REST wire form of TargetCapabilities reads back without losing data.

``TargetCapabilities`` is embedded in the ``TargetInstance`` REST response, so
``model_dump_json()`` / ``model_validate_json()`` is a real round trip: the CLI does
exactly this on every ``GET /api/targets`` payload. Serialization emits the modality
*combinations* as a sorted list of sorted lists (the ``frozenset[frozenset]`` fields
have no order of their own), so reading the wire form back has to rebuild those
combinations -- otherwise a non-text target silently reads back as text-only and the
object contradicts the payload it came from. The flattened ``supported_*_modalities``
projections stay on the wire for the UI.
"""

def test_multi_combination_profile_survives_the_wire_round_trip(self):
# More than one combination per field: reconstructing a single combination from
# the flattened projection would still be readable but would not be equal.
caps = TargetCapabilities(
input_modalities=frozenset(
{frozenset({"text"}), frozenset({"text", "image_path"}), frozenset({"audio_path", "text", "url"})}
),
output_modalities=frozenset({frozenset({"text"}), frozenset({"audio_path", "text"})}),
)

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.input_modalities == caps.input_modalities
assert restored.output_modalities == caps.output_modalities

def test_wire_form_carries_combinations_and_flattened_projections(self):
caps = TargetCapabilities(
input_modalities=frozenset({frozenset({"text"}), frozenset({"image_path", "text"})}),
)

payload = caps.model_dump(mode="json")

# Combinations, sorted within and across, so equal capabilities serialize equally.
assert payload["input_modalities"] == [["image_path", "text"], ["text"]]
# The flattened projection the UI reads stays on the wire.
assert payload["supported_input_modalities"] == ["image_path", "text"]
assert payload["supported_output_modalities"] == ["text"]

def test_serialized_ordering_is_stable_across_equivalent_objects(self):
# frozenset iteration order varies with the strings' hashes, so two objects built
# from differently-ordered inputs must still serialize identically.
first = TargetCapabilities(
input_modalities=frozenset(
{frozenset({"url"}), frozenset({"audio_path", "text", "url"}), frozenset({"text"})}
)
)
second = TargetCapabilities(
input_modalities=frozenset(
{frozenset({"text"}), frozenset({"audio_path", "text", "url"}), frozenset({"url"})}
)
)

assert first.model_dump_json() == second.model_dump_json()

def test_non_text_output_target_does_not_read_back_as_text_only(self):
caps = TargetCapabilities(output_modalities=frozenset({frozenset({"image_path"})}))

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.output_modalities == caps.output_modalities
assert restored.supported_output_modalities == ["image_path"]

def test_non_text_input_target_does_not_read_back_as_text_only(self):
caps = TargetCapabilities(input_modalities=frozenset({frozenset({"audio_path"})}))

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.input_modalities == caps.input_modalities
assert restored.supported_input_modalities == ["audio_path"]

def test_known_model_capabilities_survive_the_wire_round_trip(self):
caps = get_known_capabilities("gpt-4o")
assert caps is not None

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.supported_input_modalities == caps.supported_input_modalities
assert restored.supported_output_modalities == caps.supported_output_modalities
for field in ("supports_multi_turn", "supports_system_prompt", "supports_json_output"):
assert getattr(restored, field) == getattr(caps, field)

def test_capability_helpers_agree_after_the_wire_round_trip(self):
caps = TargetCapabilities(
supports_multi_turn=True,
output_modalities=frozenset({frozenset({"audio_path"})}),
)

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.includes(capability=CapabilityName.MULTI_TURN) is True
assert "audio_path" in restored.supported_output_modalities

def test_empty_modalities_survive_the_wire_round_trip(self):
caps = TargetCapabilities(input_modalities=frozenset())

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps
assert restored.supported_input_modalities == []
assert restored.input_modalities == frozenset()

def test_default_capabilities_are_unchanged_by_the_round_trip(self):
caps = TargetCapabilities()

restored = TargetCapabilities.model_validate_json(caps.model_dump_json())

assert restored == caps

def test_flattened_projection_in_the_payload_is_not_read_back_as_state(self):
# ``supported_*_modalities`` is a derived projection, so a payload carrying a stale
# one must not be able to override the combinations it is derived from.
caps = TargetCapabilities.model_validate(
{
"input_modalities": [["image_path"]],
"supported_input_modalities": ["text"],
}
)

assert caps.input_modalities == frozenset({frozenset({"image_path"})})
assert caps.supported_input_modalities == ["image_path"]

def test_nested_target_instance_round_trip_preserves_capabilities(self):
# ``TargetInstance`` is what the REST layer actually serves and what the CLI
# validates, so the capabilities have to survive that nesting -- including the
# inner targets of a composite target.
inner = TargetInstance(
target_registry_name="openai_chat",
identifier=TargetIdentifier(class_name="OpenAIChatTarget", class_module="pyrit.prompt_target.openai"),
capabilities=TargetCapabilities(
input_modalities=frozenset({frozenset({"text"}), frozenset({"image_path", "text"})}),
output_modalities=frozenset({frozenset({"text"})}),
),
)
composite = TargetInstance(
target_registry_name="round_robin",
identifier=TargetIdentifier(class_name="RoundRobinTarget", class_module="pyrit.prompt_target.round_robin"),
capabilities=TargetCapabilities(
supports_multi_turn=True,
input_modalities=frozenset({frozenset({"audio_path", "text"})}),
output_modalities=frozenset({frozenset({"audio_path"}), frozenset({"text"})}),
),
inner_targets=[inner],
)

restored = TargetInstance.model_validate_json(composite.model_dump_json())

assert restored == composite
assert restored.capabilities.input_modalities == composite.capabilities.input_modalities
assert restored.capabilities.output_modalities == composite.capabilities.output_modalities
assert restored.capabilities.supported_input_modalities == ["audio_path", "text"]
assert restored.inner_targets is not None
assert restored.inner_targets[0] == inner
assert restored.inner_targets[0].capabilities.input_modalities == inner.capabilities.input_modalities

def test_non_mapping_payload_is_left_to_pydantic(self):
with pytest.raises(ValidationError):
TargetCapabilities.model_validate_json('"not an object"')
Loading