From 133e4f1b218447824eba0d51edf4ee210d310e18 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 1 Oct 2026 20:03:36 +0000 Subject: [PATCH 1/8] FIX: Validate target and converter parameter types before construction The registry coerced string values and enums, but any other JSON value was passed to the constructor unchecked. A list for an integer or an object for a float either built a component with invalid state or failed inside the constructor as a server error. resolve_constructor_args now checks JSON values (null, strings, numbers, booleans, arrays, and objects) against the declared parameter type with strict pydantic JSON validation before construction, and raises ValueError on a mismatch. Abstract collection types such as Collection, Sequence, and Iterable are checked as arrays, and protocol types require a value that provides the protocol's members. Values that would need conversion into objects such as paths, enums, models, or other classes are rejected too, since the constructor would receive raw JSON. Lists of enum or literal choices are converted through Parameter.coerce_value like a single enum, so they accept the choices the catalog advertises. Arrays become tuples or sets where the type declares them; other values are passed through as received, so integers still satisfy floats. Live Python objects (including list and dict subclasses) and annotations pydantic cannot resolve (such as forward references) are left to the constructor, so in-process callers are unchanged. Registry reference parameters such as converter_target and targets only accept registry names or instances, but a JSON number, boolean, or object was passed to the constructor as-is. These now raise ValueError too. WordDocConverter and A2ATarget imported some annotation types only for type checking, so the registry could not see existing_docx as a Path parameter (skipping the upload handling for it) or check prompt_template and auth_token. Those types are now imported at runtime. --- pyrit/converter/word_doc_converter.py | 6 +- pyrit/prompt_target/a2a_target.py | 4 +- pyrit/registry/resolution.py | 222 ++++++++++++- tests/unit/backend/test_converter_service.py | 10 + .../unit/converter/test_word_doc_converter.py | 13 + tests/unit/registry/test_resolution.py | 291 +++++++++++++++++- .../score/test_garak_exploitation_scorer.py | 7 +- 7 files changed, 529 insertions(+), 24 deletions(-) diff --git a/pyrit/converter/word_doc_converter.py b/pyrit/converter/word_doc_converter.py index aad439182e..02e109a6e2 100644 --- a/pyrit/converter/word_doc_converter.py +++ b/pyrit/converter/word_doc_converter.py @@ -7,6 +7,7 @@ import hashlib from dataclasses import dataclass from io import BytesIO +from pathlib import Path # noqa: TC003 - registry annotation resolution from typing import TYPE_CHECKING, Any from docx import Document @@ -14,12 +15,11 @@ from pyrit.common.logger import logger from pyrit.converter.converter import Converter, ConverterResult from pyrit.memory import data_serializer_factory +from pyrit.models import SeedPrompt # noqa: TC001 - registry annotation resolution if TYPE_CHECKING: - from pathlib import Path - from pyrit.memory import DataTypeSerializer - from pyrit.models import ComponentIdentifier, PromptDataType, SeedPrompt + from pyrit.models import ComponentIdentifier, PromptDataType @dataclass diff --git a/pyrit/prompt_target/a2a_target.py b/pyrit/prompt_target/a2a_target.py index a434e9e9bb..b607f701a3 100644 --- a/pyrit/prompt_target/a2a_target.py +++ b/pyrit/prompt_target/a2a_target.py @@ -8,7 +8,7 @@ import math import time import uuid -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Awaitable, Callable # noqa: TC003 - registry annotation resolution from dataclasses import dataclass from email.utils import parsedate_to_datetime from typing import TYPE_CHECKING, Any, Literal, cast @@ -31,8 +31,6 @@ from pyrit.prompt_target.common.utils import limit_requests_per_minute if TYPE_CHECKING: - from collections.abc import Awaitable, Callable - from a2a.client import Client from a2a.types import a2a_pb2 diff --git a/pyrit/registry/resolution.py b/pyrit/registry/resolution.py index 3729b23af0..284eb6d8bb 100644 --- a/pyrit/registry/resolution.py +++ b/pyrit/registry/resolution.py @@ -37,15 +37,39 @@ from __future__ import annotations import copy +import functools import inspect +import json import logging +import operator import re import types -from collections.abc import Collection, Sequence +from collections.abc import Collection, Iterable, Sequence from enum import Enum -from typing import TYPE_CHECKING, Any, Protocol, TypeAlias, Union, get_args, get_origin, get_type_hints - -from pydantic import TypeAdapter, ValidationError +from typing import ( + TYPE_CHECKING, + Annotated, + Any, + Literal, + Protocol, + TypeAlias, + Union, + get_args, + get_origin, + get_type_hints, +) + +from pydantic import ( + AfterValidator, + ConfigDict, + PydanticSchemaGenerationError, + PydanticUndefinedAnnotation, + PydanticUserError, + TypeAdapter, + ValidationError, +) +from pydantic_core import SchemaError +from typing_extensions import get_protocol_members, is_protocol from pyrit.common.apply_defaults import REQUIRED_VALUE, _RequiredValueSentinel from pyrit.common.brick_contract import init_parameters_are_forwarded @@ -409,7 +433,8 @@ def _resolve_single_reference( Resolve a single registry-reference value to a stored instance. A string value is looked up by name in the paired registry. An already-built - instance passes through unchanged. + instance passes through unchanged. Other JSON data (a number or an object) can + never name an instance, so it is rejected instead of reaching the constructor. Args: value (Any): The raw value (a registry name, or an instance to pass through). @@ -421,9 +446,15 @@ def _resolve_single_reference( Any: The resolved instance. Raises: - ValueError: If the name is not registered. + ValueError: If the name is not registered, or the value is JSON data rather than a name + or instance. """ if not isinstance(value, str): + if value is not None and _is_json_data(value): + raise ValueError( + f"{owner}.{name}: expected a registry name or instance for this reference, " + f"but got {type(value).__name__}." + ) return value registry = getter() @@ -478,8 +509,9 @@ def _resolve_registry_reference( Any: The resolved instance, or a list of resolved instances. Raises: - ValueError: If a name is not registered, or the value's shape (list vs. - scalar) does not match the reference's arity. + ValueError: If a name is not registered, a value is JSON data rather than a + name or instance, or the value's shape (list vs. scalar) does not match the + reference's arity. """ if get_origin(annotation) is list: if not isinstance(value, list): @@ -522,7 +554,8 @@ def resolve_reference_value( Any: The resolved instance, or the value unchanged when already an instance. Raises: - ValueError: If no registry is wired for ``component_type``, or the name is not registered. + ValueError: If no registry is wired for ``component_type``, the name is not registered, + or the value is JSON data rather than a name or instance. """ getter = _registry_getter_for_component_type(component_type) if getter is None: @@ -541,8 +574,10 @@ def resolve_constructor_args( Derives the ``Parameter`` contract for ``cls`` and applies it to ``raw_args``. For each raw argument: validate it is a declared parameter; - resolve registry-reference parameters by name; coerce simple string values - via ``Parameter.coerce_value``; pass everything else through unchanged. + resolve registry-reference parameters by name; coerce simple string values, + enums, and JSON lists of enum or literal choices via ``Parameter.coerce_value``; + check other JSON values against the declared type; pass live Python objects + through unchanged. Args: cls (type): The class being built. @@ -556,7 +591,8 @@ def resolve_constructor_args( Raises: ValueError: If an argument is not a declared parameter, a registry - reference cannot be resolved, or a simple value cannot be coerced. + reference cannot be resolved, a simple value cannot be coerced, or a + JSON value does not match the declared type. """ by_name = {param.name: param for param in derive_parameters(cls=cls, identifier_type=identifier_type)} @@ -585,19 +621,175 @@ def resolve_constructor_args( ) elif param.variants is not None: resolved[name] = _resolve_structured_input(parameter=param, value=value) - elif (isinstance(value, str) and param.is_string_coercible) or ( - isinstance(value_type, type) and issubclass(value_type, Enum) + elif ( + (isinstance(value, str) and param.is_string_coercible) + or (isinstance(value_type, type) and issubclass(value_type, Enum)) + or (_is_choice_list(value_type) and _is_json_data(value)) ): try: resolved[name] = param.coerce_value(value) except (ValueError, TypeError) as e: raise ValueError(f"Parameter '{name}' of '{cls.__name__}': {e}") from e else: - resolved[name] = value + resolved[name] = _check_json_value(parameter=param, value=value, owner=cls.__name__) return resolved +def _is_choice_list(annotation: Any) -> bool: + """ + Return whether an annotation is a list of enum or literal choices. + + ``Parameter.coerce_value`` converts such a list from its JSON choice values, as it + does a single enum. + + Returns: + bool: True for ``list[SomeEnum]`` and ``list[Literal[...]]``. + """ + if get_origin(annotation) is not list: + return False + args = get_args(annotation) + element = args[0] if args else None + return get_origin(element) is Literal or (isinstance(element, type) and issubclass(element, Enum)) + + +_PLAIN_SEQUENCE_TYPES = (list, tuple, set, frozenset) + + +def _is_json_data(value: Any, *, sequence_types: tuple[type, ...] = (list,)) -> bool: + """ + Return whether a value holds only JSON-style data rather than live Python objects. + + Args: + value (Any): The value to inspect. + sequence_types (tuple[type, ...]): Sequence types accepted alongside string-keyed dicts. + + Returns: + bool: True for None, strings, numbers, and the given sequences or string-keyed dicts of + them (exact built-in types, so subclasses and enums count as live objects), without + cycles. + """ + pending = [value] + seen: set[int] = set() + while pending: + item = pending.pop() + if item is None or type(item) in (str, int, float, bool): + continue + if id(item) in seen or type(item) not in (dict, *sequence_types): + return False + seen.add(id(item)) + if type(item) is dict: + if not all(type(key) is str for key in item): + return False + pending.extend(item.values()) + else: + pending.extend(item) + return True + + +def _json_projection(annotation: Any) -> Any: + """ + Project an annotation onto the types its JSON form is validated against. + + Abstract collections become lists and regex patterns become strings at any depth, + including inside unions and containers, so pydantic can validate their JSON form. + Protocols become a check that the value provides every protocol member. + + Returns: + Any: The projected annotation. + """ + origin = get_origin(annotation) or annotation + if origin is re.Pattern: + return str + if is_protocol(origin): + return Annotated[object, AfterValidator(functools.partial(_require_protocol_members, protocol=origin))] + args = get_args(annotation) + if origin in (Collection, Sequence, Iterable): + if not args: + return list + element_type = _json_projection(args[0]) + return list[element_type] + if not args: + return annotation + projected = tuple(_json_projection(arg) for arg in args) + if origin in (Union, types.UnionType): + return functools.reduce(operator.or_, projected) + return origin[projected] + + +def _require_protocol_members(value: object, *, protocol: type) -> object: + """ + Require a value to provide every member of a protocol. + + Returns: + object: The unchanged value. + + Raises: + ValueError: If the value lacks a protocol member. + """ + missing = sorted(name for name in get_protocol_members(protocol) if not hasattr(value, name)) + if missing: + raise ValueError(f"{protocol.__name__} requires {', '.join(missing)}") + return value + + +def _json_adapter(annotation: Any) -> TypeAdapter[Any]: + """ + Build a JSON validator for a parameter annotation. + + Annotations naming arbitrary classes are retried as instance checks, which no JSON value + satisfies, so JSON input for an object-only parameter is rejected instead of skipped. + + Returns: + TypeAdapter[Any]: The validator. + """ + try: + return TypeAdapter(annotation) + except PydanticSchemaGenerationError: + return TypeAdapter(annotation, config=ConfigDict(arbitrary_types_allowed=True)) + + +def _check_json_value(*, parameter: Parameter, value: Any, owner: str) -> Any: + """ + Check a JSON value against the parameter's declared type before construction. + + JSON values must match the declared type, and must not need conversion into objects such + as paths, enums, or models, which the constructor would otherwise receive as raw JSON. + JSON arrays become tuples or sets where the type declares them, and protocol-typed values + must provide the protocol's members. Live Python objects from in-process callers, and + annotations pydantic cannot resolve or build (such as forward references), are left to the + constructor. + + Args: + parameter (Parameter): The declared parameter. + value (Any): The raw supplied value. + owner (str): The class name used in error messages. + + Returns: + Any: The value to pass to the constructor. + + Raises: + ValueError: If the JSON value does not match the declared type. + """ + if parameter.param_type is None or not _is_json_data(value): + return value + try: + adapter = _json_adapter(_json_projection(parameter.param_type)) + validated = adapter.validate_json(json.dumps(value), strict=True) + plain = _is_json_data(validated, sequence_types=_PLAIN_SEQUENCE_TYPES) + unchanged = plain and validated == value + except ValidationError as exc: + raise ValueError(f"Parameter '{parameter.name}' of '{owner}' expects {parameter.type_name}.") from exc + except (PydanticUserError, PydanticUndefinedAnnotation, SchemaError, RecursionError, TypeError, ValueError): + logger.debug("Skipping JSON type check for %s.%s", owner, parameter.name) + return value + if not plain: + raise ValueError( + f"Parameter '{parameter.name}' of '{owner}' expects {parameter.type_name}, which cannot be built from JSON." + ) + return value if unchanged else validated + + def _resolve_structured_input(*, parameter: Parameter, value: Any) -> Any: """ Build a declared structured-input variant from its JSON representation. diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index 26497ac54b..7d34ecc42a 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -413,6 +413,16 @@ async def test_create_converter_raises_for_invalid_type(self) -> None: with pytest.raises(ValueError, match="not found"): await service.create_converter_async(request=request) + async def test_create_converter_rejects_wrong_parameter_type(self) -> None: + """A JSON value that does not match the declared type is a validation error, not a crash.""" + service = ConverterService() + request = CreateConverterRequest(name="caesar", type="CaesarConverter", params={"caesar_offset": [1]}) + + with pytest.raises(ValueError, match="caesar_offset"): + await service.create_converter_async(request=request) + + assert service.get_converter_object(converter_id="caesar") is None + async def test_create_converter_success(self) -> None: """Test successful converter creation.""" service = ConverterService() diff --git a/tests/unit/converter/test_word_doc_converter.py b/tests/unit/converter/test_word_doc_converter.py index 1ca2d6b891..e715ae6b27 100644 --- a/tests/unit/converter/test_word_doc_converter.py +++ b/tests/unit/converter/test_word_doc_converter.py @@ -10,6 +10,7 @@ from pyrit.converter import ConverterResult, WordDocConverter from pyrit.models import SeedPrompt +from pyrit.registry.resolution import derive_parameters, resolve_constructor_args @pytest.fixture @@ -165,6 +166,18 @@ def test_constructor_custom_placeholder() -> None: assert converter._injection_config.placeholder == "<>" +def test_existing_docx_is_declared_as_path_parameter() -> None: + """The registry must see ``existing_docx`` as a ``Path`` so API callers upload the file.""" + parameter = next(param for param in derive_parameters(cls=WordDocConverter) if param.name == "existing_docx") + assert parameter.is_path + + +def test_prompt_template_json_value_is_rejected() -> None: + """The registry must resolve ``prompt_template`` so JSON values are checked before construction.""" + with pytest.raises(ValueError, match="prompt_template"): + resolve_constructor_args(cls=WordDocConverter, raw_args={"prompt_template": "hello"}) + + def test_build_identifier_without_template() -> None: """_build_identifier should return correct params when no template is set.""" converter = WordDocConverter() diff --git a/tests/unit/registry/test_resolution.py b/tests/unit/registry/test_resolution.py index 5cf32d4e85..dd70877173 100644 --- a/tests/unit/registry/test_resolution.py +++ b/tests/unit/registry/test_resolution.py @@ -5,15 +5,20 @@ Tests for the shared registry constructor-argument resolution primitive. """ +import contextlib +import json +from collections.abc import Collection +from dataclasses import dataclass from enum import Enum -from typing import Literal +from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal, Protocol import pytest from pyrit.common import REQUIRED_VALUE, forward_init_parameters from pyrit.common.apply_defaults import _RequiredValueSentinel from pyrit.models import Message, MessagePiece -from pyrit.models.identifiers import ConverterIdentifier, TargetIdentifier +from pyrit.models.identifiers import ConverterIdentifier, ScorerIdentifier, TargetIdentifier from pyrit.models.parameter import ComponentType from pyrit.prompt_target import PromptTarget from pyrit.registry.components import ConverterRegistry, ScorerRegistry, TargetRegistry @@ -24,6 +29,9 @@ resolve_constructor_args, ) +if TYPE_CHECKING: + from pyrit.prompt_target import PromptTarget as _TypeCheckingOnlyTarget + class MockPromptTarget(PromptTarget): """Minimal PromptTarget for registry-resolution tests.""" @@ -138,6 +146,79 @@ def __init__(self, *, targets: list[PromptTarget]) -> None: self.targets = targets +@dataclass +class _Settings: + level: int = 0 + + +class _Provider(Protocol): + def provide(self) -> str: ... + + +class _Sized(Protocol): + def __len__(self) -> int: ... + + +class _Unresolved: + """Helper whose annotation names a type-checking-only import, as many components do.""" + + def __init__(self, *, target: "_TypeCheckingOnlyTarget | None" = None) -> None: + self.target = target + + +class _Handle: + """A live object type that no JSON value can represent.""" + + +class _Bag(list[int]): + """A live list subclass that callers pass as an existing object.""" + + +class _JsonShaped: + """Helper whose constructor takes container and object parameters that JSON callers supply.""" + + def __init__( + self, + *, + color: tuple[int, int, int] = (0, 0, 0), + weights: list[int] | None = None, + extra: dict[str, int] | None = None, + speed: _Speed | None = None, + speeds: list[_Speed] | None = None, + modes: list[Literal["a", "b"] | None] | None = None, + note: str | None = None, + words: Collection[str] | None = None, + settings: _Settings | None = None, + location: _Speed | Path | None = None, + provider: _Provider | None = None, + sized: _Sized | None = None, + options: dict[str, Any] | None = None, + handle: _Handle | None = None, + groups: dict[str, Collection[str]] | None = None, + choices: Collection[str] | _Handle | None = None, + bag: _Bag | None = None, + anything: Collection | None = None, + ) -> None: + self.color = color + self.weights = weights + self.extra = extra + self.speed = speed + self.speeds = speeds + self.modes = modes + self.note = note + self.words = words + self.settings = settings + self.location = location + self.provider = provider + self.sized = sized + self.options = options + self.handle = handle + self.groups = groups + self.choices = choices + self.bag = bag + self.anything = anything + + def _resolve(cls: type, raw_args: dict[str, object], *, identifier_type: type | None = None) -> dict[str, object]: """Resolve ``raw_args`` against the derived parameter contract for ``cls``.""" return resolve_constructor_args(cls=cls, raw_args=raw_args, identifier_type=identifier_type) @@ -259,6 +340,212 @@ def test_unknown_registry_reference_empty_registry_hint(self, empty_target_regis with pytest.raises(ValueError, match="is empty"): _resolve(_NeedsTarget, {"converter_target": "missing"}, identifier_type=ConverterIdentifier) + @pytest.mark.parametrize( + ("cls", "raw_args"), + [ + (_SimpleOnly, {"count": {"value": 1}}), + (_SimpleOnly, {"count": [5]}), + (_SimpleOnly, {"count": True}), + (_SimpleOnly, {"count": 1.5}), + (_SimpleOnly, {"count": None}), + (_SimpleOnly, {"ratio": {"value": 1}}), + (_SimpleOnly, {"flag": 1}), + (_JsonShaped, {"note": 5}), + (_JsonShaped, {"weights": [1, "2"]}), + (_JsonShaped, {"extra": "not-an-object"}), + (_JsonShaped, {"color": [1, 2]}), + (_JsonShaped, {"color": None}), + (_JsonShaped, {"words": [1, "a"]}), + (_JsonShaped, {"words": "the"}), + (_JsonShaped, {"groups": {"group": [1]}}), + (_JsonShaped, {"choices": {}}), + (_JsonShaped, {"anything": 7}), + (_JsonShaped, {"modes": ["c"]}), + ], + ) + def test_rejects_json_value_of_wrong_type(self, cls: type, raw_args: dict[str, object]) -> None: + with pytest.raises(ValueError, match="expects"): + _resolve(cls, raw_args) + + @pytest.mark.parametrize( + ("cls", "raw_args"), + [ + (_SimpleOnly, {"ratio": 1}), + (_SimpleOnly, {"count": 3, "flag": False}), + (_JsonShaped, {"weights": [1, 2], "extra": {"a": 1}}), + (_JsonShaped, {"note": None, "settings": None}), + (_JsonShaped, {"words": ["the", "a"]}), + (_JsonShaped, {"groups": {"group": ["one", "two"]}}), + (_JsonShaped, {"choices": ["the", "a"]}), + (_JsonShaped, {"anything": [1, "a"]}), + (_JsonShaped, {"modes": []}), + (_JsonShaped, {"modes": ["a", None]}), + ], + ) + def test_accepts_matching_json_value_unchanged(self, cls: type, raw_args: dict[str, object]) -> None: + assert _resolve(cls, raw_args) == raw_args + + def test_json_array_becomes_tuple_for_tuple_parameter(self) -> None: + assert _resolve(_JsonShaped, {"color": [10, 20, 30]})["color"] == (10, 20, 30) + + @pytest.mark.parametrize("raw_args", [{"settings": {"level": 1}}, {"location": "fast"}, {"location": "/tmp/x"}]) + def test_rejects_json_value_that_needs_an_object(self, raw_args: dict[str, object]) -> None: + with pytest.raises(ValueError, match="cannot be built from JSON"): + _resolve(_JsonShaped, raw_args) + + def test_live_objects_pass_through_unchecked(self) -> None: + color = object() + weights = [object()] + options = {"timeout": (5.0, 10.0)} + extra: dict[str, object] = {} + extra["self"] = extra + bag = _Bag([1, 2]) + live = {"color": color, "weights": weights, "options": options, "extra": extra, "bag": bag} + + resolved = _resolve(_JsonShaped, live) + + assert all(resolved[name] is value for name, value in live.items()) + + def test_deeply_nested_json_value_is_rejected(self) -> None: + nested: object = 1 + for _ in range(500): + nested = [nested] + + with pytest.raises(ValueError, match="expects"): + _resolve(_SimpleOnly, {"count": nested}) + + def test_protocol_parameter_requires_protocol_members(self) -> None: + with pytest.raises(ValueError, match="provider"): + _resolve(_JsonShaped, {"provider": "anything"}) + with pytest.raises(ValueError, match="sized"): + _resolve(_JsonShaped, {"sized": 7}) + + assert _resolve(_JsonShaped, {"provider": None, "sized": [1, 2]}) == {"provider": None, "sized": [1, 2]} + + def test_unresolved_annotation_is_left_to_constructor(self) -> None: + assert _resolve(_Unresolved, {"target": {"a": 1}}) == {"target": {"a": 1}} + + def test_json_value_for_object_parameter_is_rejected(self) -> None: + handle = _Handle() + + with pytest.raises(ValueError, match="handle"): + _resolve(_JsonShaped, {"handle": {}}) + assert _resolve(_JsonShaped, {"handle": None}) == {"handle": None} + assert _resolve(_JsonShaped, {"handle": handle})["handle"] is handle + + @pytest.mark.parametrize( + ("registry_type", "identifier_type", "type_name", "raw_args"), + [ + (TargetRegistry, TargetIdentifier, "TextTarget", {"custom_configuration": {}}), + (TargetRegistry, TargetIdentifier, "A2ATarget", {"auth_token": {"token": "x"}}), + (TargetRegistry, TargetIdentifier, "OpenAIResponseTarget", {"tool_providers": [{"name": "x"}]}), + (ConverterRegistry, ConverterIdentifier, "TokenBijectionConverter", {"tokenizer": "name"}), + ], + ) + def test_registered_component_rejects_json_for_object_parameter( + self, registry_type: type, identifier_type: type, type_name: str, raw_args: dict[str, object] + ) -> None: + cls = registry_type.get_registry_singleton().get_class(type_name) + + with pytest.raises(ValueError, match="expects"): + _resolve(cls, raw_args, identifier_type=identifier_type) + + @pytest.mark.parametrize( + ("registry_type", "identifier_type", "type_name", "raw_args"), + [ + (ConverterRegistry, ConverterIdentifier, "FlipConverter", {"converter_target": {"name": "x"}}), + (TargetRegistry, TargetIdentifier, "RoundRobinTarget", {"targets": [{"name": "x"}]}), + ], + ) + def test_registry_reference_rejects_json_object( + self, registry_type: type, identifier_type: type, type_name: str, raw_args: dict[str, object] + ) -> None: + cls = registry_type.get_registry_singleton().get_class(type_name) + + with pytest.raises(ValueError, match="registry name or instance"): + _resolve(cls, raw_args, identifier_type=identifier_type) + + def test_registered_target_accepts_string_token(self) -> None: + cls = TargetRegistry.get_registry_singleton().get_class("A2ATarget") + + assert _resolve(cls, {"auth_token": "token"}, identifier_type=TargetIdentifier) == {"auth_token": "token"} + + @pytest.mark.parametrize( + ("type_name", "raw_args", "expected"), + [ + ("ImageCompressionConverter", {"background_color": [10, 20, 30]}, {"background_color": (10, 20, 30)}), + ("SATAMaskingConverter", {"stopwords": ["the", "a"]}, {"stopwords": ["the", "a"]}), + ], + ) + def test_registered_converter_json_values( + self, type_name: str, raw_args: dict[str, object], expected: dict[str, object] + ) -> None: + cls = ConverterRegistry.get_registry_singleton().get_class(type_name) + + assert _resolve(cls, raw_args, identifier_type=ConverterIdentifier) == expected + + @pytest.mark.parametrize( + ("registry_type", "identifier_type"), + [(ConverterRegistry, ConverterIdentifier), (TargetRegistry, TargetIdentifier)], + ) + def test_registered_parameters_only_raise_value_error(self, registry_type: type, identifier_type: type) -> None: + registry = registry_type.get_registry_singleton() + for type_name in registry.get_class_names(): + cls = registry.get_class(type_name) + for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): + for value in (None, 1, 1.5, True, "text", [], [1, "a"], {"a": [1]}): + with contextlib.suppress(ValueError): + _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) + + @pytest.mark.parametrize( + ("registry_type", "identifier_type"), + [(ConverterRegistry, ConverterIdentifier), (TargetRegistry, TargetIdentifier)], + ) + def test_registered_json_defaults_are_accepted(self, registry_type: type, identifier_type: type) -> None: + registry = registry_type.get_registry_singleton() + for type_name in registry.get_class_names(): + cls = registry.get_class(type_name) + for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): + if parameter.reference is not None or parameter.default is None: + continue + try: + value = json.loads(json.dumps(parameter.default)) + except (TypeError, ValueError): + continue + _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) + + @pytest.mark.parametrize( + ("registry_type", "identifier_type"), + [ + (ConverterRegistry, ConverterIdentifier), + (TargetRegistry, TargetIdentifier), + (ScorerRegistry, ScorerIdentifier), + ], + ) + def test_registered_choices_are_accepted(self, registry_type: type, identifier_type: type) -> None: + registry = registry_type.get_registry_singleton() + for type_name in registry.get_class_names(): + cls = registry.get_class(type_name) + for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): + if parameter.reference is not None or not parameter.choices: + continue + value = [parameter.choices[0]] if parameter.is_list else parameter.choices[0] + _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) + + def test_enum_list_is_coerced_from_json_choices(self) -> None: + live = [_Speed.SLOW] + + assert _resolve(_JsonShaped, {"speeds": ["fast"]}) == {"speeds": [_Speed.FAST]} + assert _resolve(_JsonShaped, {"speeds": live})["speeds"] is live + with pytest.raises(ValueError, match="speeds"): + _resolve(_JsonShaped, {"speeds": ["bogus"]}) + + def test_collection_parameter_rejects_wrong_member_type(self) -> None: + cls = ConverterRegistry.get_registry_singleton().get_class("SATAMaskingConverter") + + with pytest.raises(ValueError, match="stopwords"): + _resolve(cls, {"stopwords": [1, "a"]}, identifier_type=ConverterIdentifier) + class TestDeriveParameters: """Tests for deriving the Parameter contract from a constructor signature.""" diff --git a/tests/unit/score/test_garak_exploitation_scorer.py b/tests/unit/score/test_garak_exploitation_scorer.py index 48738c8123..5b9e2da013 100644 --- a/tests/unit/score/test_garak_exploitation_scorer.py +++ b/tests/unit/score/test_garak_exploitation_scorer.py @@ -78,11 +78,16 @@ async def test_default_corpus_is_ready_to_score_async( assert (await scorer.score_text_async(text))[0].get_value() is True assert (await scorer.score_text_async("I cannot help with that request."))[0].get_value() is False - @pytest.mark.parametrize("payloads", [[], [""], "payload()", b"payload()", ("payload()",), [1], ["valid", None]]) + @pytest.mark.parametrize("payloads", [[], [""], b"payload()", ("payload()",)]) def test_invalid_payloads_raise(self, payloads: object) -> None: with pytest.raises(ValueError, match="nonempty list of nonempty strings"): ScorerRegistry().create_instance("GarakExploitationScorer", detector="jinja", payloads=payloads) + @pytest.mark.parametrize("payloads", ["payload()", [1], ["valid", None]]) + def test_mistyped_json_payloads_raise(self, payloads: object) -> None: + with pytest.raises(ValueError, match=r"expects list\[str\]"): + ScorerRegistry().create_instance("GarakExploitationScorer", detector="jinja", payloads=payloads) + async def test_registry_list_contract_async(self) -> None: registry = ScorerRegistry() metadata = registry.get_class_metadata(GarakExploitationScorer) From b33d822d81fe0c40363954fd2878a893f8ae2c65 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Fri, 2 Oct 2026 20:50:34 +0000 Subject: [PATCH 2/8] Replace JSON type projection with an explicit registry input contract Each derived Parameter now reports an input_kind: scalar, collection, reference, structured, in_process_only, or unsupported. The resolver converts JSON only for scalar and collection parameters, using explicit rules instead of a pydantic projection of the annotation, and rejects JSON for in_process_only and unsupported parameters. Live Python objects still pass through to constructors. Registry metadata and the converter and target catalogs report whether a class is constructible, meaning every required parameter can be supplied as JSON. framework.md documents the contract. --- doc/code/framework.md | 10 + pyrit/backend/models/converters.py | 3 + pyrit/backend/models/targets.py | 3 + pyrit/backend/services/converter_service.py | 8 +- pyrit/backend/services/target_service.py | 8 +- pyrit/models/parameter.py | 328 +++++++++++++++++- pyrit/registry/registry_metadata.py | 5 + pyrit/registry/resolution.py | 196 ++--------- tests/unit/backend/test_converter_service.py | 14 +- tests/unit/backend/test_target_service.py | 10 + tests/unit/models/test_parameter.py | 156 ++++++++- tests/unit/registry/test_registry_metadata.py | 22 ++ tests/unit/registry/test_resolution.py | 72 ++-- 13 files changed, 626 insertions(+), 209 deletions(-) diff --git a/doc/code/framework.md b/doc/code/framework.md index 09313f9db0..a0c26a8a23 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -394,6 +394,16 @@ See [message normalizers](./targets/11_message_normalizer) for capability behavi - If you are creating a component with user input (e.g. via config, REST, or automatically) it should always use the registry - If you are storing an instance of a component, it should always use the registry +- Constructor parameters are the registry's input contract. Each derived `Parameter` has an `input_kind` that says how callers supply it: + - `scalar`: `str`, `int`, `float`, `bool`, `Path`, `Path | str`, `Literal[...]`, or an `Enum`, optionally nullable. JSON callers send one string, number, or boolean; choices match by value or name. + - `collection`: a list, tuple, set, `Sequence`, `Collection`, or string-keyed `dict` / `Mapping` of scalar or collection values, never `Path`. JSON callers send an array or object; arrays become the declared tuple or set. + - `reference`: a parameter the component's identifier marks as a registry reference. Callers send a registry name, or a list of names. + - `structured`: a `StructuredParameterValue`. Callers send one of its declared variants. + - `in_process_only`: any other type (a class, callable, or protocol), or a parameter marked `opaque`. JSON is rejected; Python callers pass a live object. + - `unsupported`: no annotation, `Any`, or an annotation that cannot be resolved at runtime. JSON is rejected; resolve the annotation to make the parameter configurable. +- A union takes JSON when one of its members does (e.g. `SeedPrompt | str` takes a string), and a JSON value becomes the first member it matches. A `Path` inside a wider union or a collection is not JSON input, so JSON cannot name server files. +- The resolver enforces this contract for every caller: JSON for an `in_process_only` or `unsupported` parameter is rejected, not passed to the constructor. Live Python objects pass through unchanged, and constructors still own component-specific validation. +- A component with a required parameter that is not JSON input is not `constructible`; registry metadata and the REST catalog report this so clients do not offer it. ## [Setup](./setup/0_setup) diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 316cedb42b..057a3ff2fa 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -44,6 +44,9 @@ class ConverterTypeEntry(BaseModel): parameters: list[Parameter] = Field( default_factory=list, description="Constructor parameters for dynamic form generation" ) + constructible: bool = Field( + True, description="Whether every required parameter can be supplied through this API (see input_kind)" + ) is_llm_based: bool = Field(False, description="Whether this converter requires an LLM target") description: str | None = Field(None, description="Short description of the converter from its docstring") diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index ba9759e767..9479d042b9 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -36,6 +36,9 @@ class TargetTypeEntry(BaseModel): default_factory=list, description="Constructor parameters for dynamic form generation", ) + constructible: bool = Field( + True, description="Whether every required parameter can be supplied through this API (see input_kind)" + ) supported_auth_modes: list[Literal["api_key", "identity"]] = Field( default_factory=_default_auth_modes, description="Authentication modes this target type supports", diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index e85f91be47..3e060bec67 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -112,9 +112,10 @@ async def list_converter_types_async(self) -> ConverterTypeResponse: """ List all available converter types from the converter class registry. - Returns every constructible converter. Deciding which entries to surface - to a user is a presentation concern owned by the caller (e.g. the - frontend), not this service. + Returns every registered converter type. ``constructible`` is False when a + required parameter cannot be supplied through the API. Deciding which + entries to surface to a user is a presentation concern owned by the caller + (e.g. the frontend), not this service. Returns: ConverterTypeResponse containing all available converter classes. @@ -125,6 +126,7 @@ async def list_converter_types_async(self) -> ConverterTypeResponse: supported_input_types=list(metadata.supported_input_types), supported_output_types=list(metadata.supported_output_types), parameters=list(metadata.parameters), + constructible=metadata.constructible, is_llm_based=metadata.is_llm_based, description=metadata.class_description or None, ) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index d57188f2c7..8942cdd8f3 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -152,9 +152,10 @@ async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. - Returns every constructible target with its derived constructor - parameters and the auth modes it supports, all projected from the - registry's ``TargetMetadata``. Deciding which entries to surface to a + Returns every registered target type with its derived constructor + parameters, whether its required parameters can be supplied through the + API (``constructible``), and the auth modes it supports, all projected from + the registry's ``TargetMetadata``. Deciding which entries to surface to a user is a presentation concern owned by the caller (e.g. the frontend), not this service. @@ -166,6 +167,7 @@ async def list_target_types_async(self) -> TargetTypeResponse: TargetTypeEntry( target_type=metadata.class_name, parameters=list(metadata.parameters), + constructible=metadata.constructible, supported_auth_modes=self._get_supported_auth_modes(metadata.supported_auth_modes), description=metadata.class_description or None, ) diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index fea37a422b..edf170f02e 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -8,15 +8,36 @@ import copy import types from abc import ABC, abstractmethod +from collections.abc import Collection, Iterable, Mapping, MutableMapping, MutableSequence, MutableSet, Sequence +from collections.abc import Set as AbstractSet from dataclasses import dataclass from enum import Enum from pathlib import Path -from typing import Any, Literal, Union, get_args, get_origin +from typing import Any, ForwardRef, Literal, TypeAlias, Union, get_args, get_origin from pydantic import BaseModel, ConfigDict, Field, computed_field, field_serializer, model_validator from pyrit.common.apply_defaults import REQUIRED_VALUE +#: How registry callers supply a parameter; see ``Parameter.input_kind``. +ParameterInputKind: TypeAlias = Literal[ + "scalar", "collection", "reference", "structured", "in_process_only", "unsupported" +] +_JSON_INPUT_KINDS: frozenset[str] = frozenset({"scalar", "collection", "reference", "structured"}) +_SEQUENCE_ORIGINS: tuple[Any, ...] = ( + list, + tuple, + set, + frozenset, + Collection, + Sequence, + MutableSequence, + Iterable, + AbstractSet, + MutableSet, +) +_SET_ORIGINS: tuple[Any, ...] = (set, frozenset, AbstractSet, MutableSet) +_MAPPING_ORIGINS: tuple[Any, ...] = (dict, Mapping, MutableMapping) _SUPPORTED_SCALAR_TYPES: tuple[type, ...] = (str, int, float, bool, Path) _SCALAR_NAME_TO_TYPE: dict[str, type | types.UnionType] = { "Path": Path, @@ -85,7 +106,8 @@ class Parameter(BaseModel): *is* the allowed set) and drives ``coerce_value`` / ``validate``; it is **not** serialized. Serialization instead projects the type into the display fields ``type_name``, ``choices``, and ``is_list`` (plus ``required`` from the - ``REQUIRED_VALUE`` sentinel), so a consumer can rebuild a usable contract from + ``REQUIRED_VALUE`` sentinel and ``input_kind``, which says how registry callers + supply the value), so a consumer can rebuild a usable contract from the registry without the live type travelling on the wire. ``reference``, when set, marks the parameter as a registry reference: its value @@ -129,6 +151,14 @@ class Parameter(BaseModel): exclude=True, description="Where the parameter is consumed at build time; not serialized.", ) + wire_input_kind: ParameterInputKind | None = Field( + default=None, + exclude=True, + description=( + "The ``input_kind`` read from a serialized payload. A ``param_type`` rebuilt from display " + "fields cannot always express it, so the serialized value is kept. Not serialized." + ), + ) opaque: bool = Field( default=False, exclude=True, @@ -171,6 +201,7 @@ def _reconstruct_param_type_from_wire(cls, data: Any) -> Any: choices=data.get("choices"), is_list=bool(data.get("is_list")), ) + data.setdefault("wire_input_kind", data.get("input_kind")) if needs_reference: data["reference"] = RegistryReference( component_type=ComponentType(data["reference_type"]), @@ -227,6 +258,33 @@ def reference_type(self) -> str | None: """Registry component family this parameter references, or None.""" return self.reference.component_type.value if self.reference is not None else None + @computed_field + @property + def input_kind(self) -> ParameterInputKind: + """ + How registry callers supply this parameter. + + ``scalar`` and ``collection`` parameters take JSON values of their declared type, + ``reference`` parameters take registry names, and ``structured`` parameters take one of + their declared variants. ``in_process_only`` parameters take only live Python objects, + and ``unsupported`` parameters have a type the registry cannot describe; neither takes + JSON. + """ + if self.wire_input_kind is not None: + return self.wire_input_kind + if self.reference is not None: + return "reference" + if self.variants is not None: + return "structured" + if self.opaque: + return "in_process_only" + return _annotation_input_kind(self.param_type) + + @property + def is_json_configurable(self) -> bool: + """Whether registry callers can supply this parameter as JSON.""" + return self.input_kind in _JSON_INPUT_KINDS + @field_serializer("default") def _serialize_default(self, value: Any) -> str | list[str] | None: """ @@ -318,6 +376,43 @@ def coerce_value(self, raw_value: Any) -> Any: return _coerce_simple_value(param_name=self.name, annotation=param_type, raw_value=raw_value) return raw_value + def coerce_json_value(self, value: Any, *, owner: str) -> Any: + """ + Convert a JSON value to this parameter's declared type under the registry contract. + + A ``scalar`` or ``collection`` parameter takes JSON of its declared type: arrays become + the declared list, tuple, or set, and enum or literal choices become their members. A + path is accepted only where the parameter is itself a ``Path`` or ``Path | str``. Other + parameters take no JSON. ``None`` is accepted wherever the parameter allows it. + + Args: + value (Any): JSON data: None, a string, number, boolean, list, or string-keyed dict. + owner (str): The class name used in error messages. + + Returns: + Any: The value converted to the declared type. + + Raises: + ValueError: If the parameter cannot be supplied as JSON, or the value does not match + its declared type. + """ + if value is None and (self.default is None or type(None) in _union_members(self.param_type)): + return None + input_kind = self.input_kind + if input_kind == "in_process_only": + raise ValueError( + f"Parameter '{self.name}' of '{owner}' accepts only a Python object ({self.type_name}), not JSON." + ) + if input_kind not in ("scalar", "collection"): + raise ValueError( + f"Parameter '{self.name}' of '{owner}' has a type the registry does not support " + f"({self.type_name}), so it cannot be set from JSON." + ) + try: + return _convert_json_value(self.param_type, value, allow_path=True) + except _JsonMismatchError: + raise ValueError(f"Parameter '{self.name}' of '{owner}' expects {self.type_name}.") from None + def validate(self) -> None: # type: ignore[ty:invalid-method-override] """ Reject a declaration with an unsupported ``param_type``. @@ -398,6 +493,235 @@ def _is_scalar_param_type(annotation: Any) -> bool: return _is_enum_type(annotation) +class _JsonMismatchError(Exception): + """Raised when a JSON value does not match a declared type.""" + + +def _union_members(annotation: Any) -> tuple[Any, ...]: + """ + Return the members of a union annotation, or the annotation alone. + + Returns: + tuple[Any, ...]: The union's members, including ``NoneType`` when present. + """ + if get_origin(annotation) in (Union, types.UnionType): + return get_args(annotation) + return (annotation,) + + +def _is_unresolved(annotation: Any) -> bool: + """ + Return whether an annotation is, or contains, a forward reference that was never resolved. + + Returns: + bool: True when a string or ``ForwardRef`` stands in for a type. + """ + if isinstance(annotation, (str, ForwardRef)): + return True + if get_origin(annotation) is Literal: + return False + return any(_is_unresolved(arg) for arg in get_args(annotation)) + + +def _is_json_scalar_type(annotation: Any) -> bool: + """ + Return whether JSON expresses the annotation as one string, number, or boolean. + + Returns: + bool: True for ``str``, ``int``, ``float``, ``bool``, an ``Enum``, or a ``Literal`` of those values. + """ + if annotation in (str, int, float, bool) or _is_enum_type(annotation): + return True + return get_origin(annotation) is Literal and all( + arg is None or type(arg) in (str, int, float, bool) for arg in get_args(annotation) + ) + + +def _is_json_element_type(annotation: Any) -> bool: + """ + Return whether a collection element type can be expressed in JSON. + + Paths are not, so a collection cannot carry server file paths. + + Returns: + bool: True for JSON scalars, JSON collections, ``Any``, ``None``, and unions of those. + """ + if annotation is Any or annotation is type(None): + return True + if get_origin(annotation) in (Union, types.UnionType): + return all(_is_json_element_type(member) for member in get_args(annotation)) + return _is_json_scalar_type(annotation) or _is_json_collection_type(annotation) + + +def _is_json_collection_type(annotation: Any) -> bool: + """ + Return whether the annotation is a list, tuple, set, or string-keyed mapping of JSON values. + + Returns: + bool: True when a JSON array or object can express the annotation. + """ + origin = get_origin(annotation) or annotation + args = [arg for arg in get_args(annotation) if arg is not Ellipsis] + if origin in _MAPPING_ORIGINS: + key_type, value_type = args if len(args) == 2 else (str, Any) + return key_type in (str, Any) and _is_json_element_type(value_type) + if origin in _SEQUENCE_ORIGINS: + return all(_is_json_element_type(arg) for arg in args) + return False + + +def _member_input_kind(annotation: Any, *, in_union: bool) -> ParameterInputKind: + """ + Classify one member of a value parameter's annotation. + + A ``Path`` is a scalar only when it is the whole type; inside a wider union it can only be + passed as a live object. + + Returns: + ParameterInputKind: ``scalar``, ``collection``, ``in_process_only``, or ``unsupported``. + """ + if annotation is Any: + return "unsupported" + if annotation is Path: + return "in_process_only" if in_union else "scalar" + if _is_json_scalar_type(annotation): + return "scalar" + if _is_json_collection_type(annotation): + return "collection" + return "in_process_only" + + +def _annotation_input_kind(annotation: Any) -> ParameterInputKind: + """ + Classify a value parameter's annotation under the registry input contract. + + A union takes JSON when any member does; its other members can still be passed as live + objects. Unannotated parameters, ``Any``, and unresolved annotations are unsupported. + + Returns: + ParameterInputKind: ``scalar``, ``collection``, ``in_process_only``, or ``unsupported``. + """ + if annotation is None or _is_unresolved(annotation): + return "unsupported" + if _is_path_or_str(annotation): + return "scalar" + members = [member for member in _union_members(annotation) if member is not type(None)] + kinds = {_member_input_kind(member, in_union=len(members) > 1) for member in members} + precedence: tuple[ParameterInputKind, ...] = ("unsupported", "collection", "scalar", "in_process_only") + return next((kind for kind in precedence if kind in kinds), "unsupported") + + +def _convert_json_value(annotation: Any, value: Any, *, allow_path: bool) -> Any: + """ + Convert a JSON value to ``annotation``. + + A union takes the value as its first member that matches. ``allow_path`` is True only + for the parameter's own type, so paths are never read from inside a union or collection, + except for an explicit ``Path | str``. + + Returns: + Any: The converted value. + + Raises: + _JsonMismatchError: If the value does not match the annotation. + """ + if annotation is Any: + return value + all_members = _union_members(annotation) + members = [member for member in all_members if member is not type(None)] + if value is None: + if len(members) < len(all_members): + return None + raise _JsonMismatchError + if allow_path and _is_path_or_str(annotation): + if type(value) is str: + return value + raise _JsonMismatchError + if len(members) == 1: + return _convert_json_member(members[0], value, allow_path=allow_path) + for member in members: + try: + return _convert_json_member(member, value, allow_path=False) + except _JsonMismatchError: + continue + raise _JsonMismatchError + + +def _convert_json_member(annotation: Any, value: Any, *, allow_path: bool) -> Any: + """ + Convert a non-null JSON value to one non-union type. + + Choices match the way ``coerce_value`` matches them, so their display strings are accepted. + + Returns: + Any: The converted value. + + Raises: + _JsonMismatchError: If the value does not match the type. + """ + if annotation is Any: + return value + if annotation in (str, bool): + if type(value) is annotation: + return value + elif annotation is int: + if type(value) is int: + return value + elif annotation is float: + if type(value) in (int, float): + try: + return float(value) + except OverflowError: + pass + elif annotation is Path: + if allow_path and type(value) is str: + return Path(value) + elif get_origin(annotation) is Literal or _is_enum_type(annotation): + if type(value) in (str, int, float, bool): + try: + return _coerce_simple_value(param_name="", annotation=annotation, raw_value=value) + except ValueError: + pass + else: + return _convert_json_collection(annotation, value) + raise _JsonMismatchError + + +def _convert_json_collection(annotation: Any, value: Any) -> Any: + """ + Convert a JSON array or object to a declared collection type. + + Returns: + Any: A list, tuple, set, frozenset, or dict holding the converted items. + + Raises: + _JsonMismatchError: If the value does not match the collection type. + """ + origin = get_origin(annotation) or annotation + args = get_args(annotation) + if origin in _MAPPING_ORIGINS: + key_type, value_type = args if len(args) == 2 else (str, Any) + if type(value) is not dict or key_type not in (str, Any): + raise _JsonMismatchError + return {key: _convert_json_value(value_type, item, allow_path=False) for key, item in value.items()} + if origin not in _SEQUENCE_ORIGINS or type(value) is not list: + raise _JsonMismatchError + if origin is tuple and args and not (len(args) == 2 and args[1] is Ellipsis): + if len(args) != len(value): + raise _JsonMismatchError + return tuple(_convert_json_value(arg, item, allow_path=False) for arg, item in zip(args, value, strict=True)) + element_type = args[0] if args else Any + items = [_convert_json_value(element_type, item, allow_path=False) for item in value] + if origin is tuple: + return tuple(items) + if origin in _SET_ORIGINS: + try: + return frozenset(items) if origin is frozenset else set(items) + except TypeError as exc: + raise _JsonMismatchError from exc + return items + + def _coerce_simple_value(*, param_name: str, annotation: Any, raw_value: Any) -> Any: """ Coerce ``raw_value`` to a scalar ``annotation`` — the shared coercion core. diff --git a/pyrit/registry/registry_metadata.py b/pyrit/registry/registry_metadata.py index 24d21fd5fb..8a6b0c0369 100644 --- a/pyrit/registry/registry_metadata.py +++ b/pyrit/registry/registry_metadata.py @@ -51,6 +51,11 @@ class RegistryMetadata: parameters: tuple[Parameter, ...] = field(kw_only=True, default=()) class_attributes: Mapping[str, Any] = field(kw_only=True, default_factory=dict) + @property + def constructible(self) -> bool: + """Whether registry callers can supply every required parameter as JSON (see ``Parameter.input_kind``).""" + return all(parameter.is_json_configurable for parameter in self.parameters if parameter.required) + @staticmethod def description_from_docstring(cls: type, *, fallback: str = "") -> str: """ diff --git a/pyrit/registry/resolution.py b/pyrit/registry/resolution.py index 284eb6d8bb..42992f68ce 100644 --- a/pyrit/registry/resolution.py +++ b/pyrit/registry/resolution.py @@ -19,8 +19,10 @@ - **Resolve from a constructor** (``resolve_constructor_args``): derive the contract for a class and turn a flat dict of raw arguments into constructor-ready keyword arguments — coercing simple string values via - ``Parameter.coerce_value`` and resolving registry-reference parameters by name - from the owning domain's registry. Defaults are left to the constructor. + ``Parameter.coerce_value``, converting other JSON values under the registry + input contract (``Parameter.input_kind`` / ``Parameter.coerce_json_value``), + and resolving registry-reference parameters by name from the owning domain's + registry. Defaults are left to the constructor. - **Resolve from a declared list** (``resolve_declared_params``): the sibling for a component that declares an explicit ``list[Parameter]`` (e.g. a scenario's ``supported_parameters()``). It has no references, coerces every supplied @@ -37,39 +39,14 @@ from __future__ import annotations import copy -import functools import inspect -import json import logging -import operator import re import types -from collections.abc import Collection, Iterable, Sequence -from enum import Enum -from typing import ( - TYPE_CHECKING, - Annotated, - Any, - Literal, - Protocol, - TypeAlias, - Union, - get_args, - get_origin, - get_type_hints, -) - -from pydantic import ( - AfterValidator, - ConfigDict, - PydanticSchemaGenerationError, - PydanticUndefinedAnnotation, - PydanticUserError, - TypeAdapter, - ValidationError, -) -from pydantic_core import SchemaError -from typing_extensions import get_protocol_members, is_protocol +from collections.abc import Collection, Sequence +from typing import TYPE_CHECKING, Any, Protocol, TypeAlias, Union, get_args, get_origin, get_type_hints + +from pydantic import TypeAdapter, ValidationError from pyrit.common.apply_defaults import REQUIRED_VALUE, _RequiredValueSentinel from pyrit.common.brick_contract import init_parameters_are_forwarded @@ -574,10 +551,11 @@ def resolve_constructor_args( Derives the ``Parameter`` contract for ``cls`` and applies it to ``raw_args``. For each raw argument: validate it is a declared parameter; - resolve registry-reference parameters by name; coerce simple string values, - enums, and JSON lists of enum or literal choices via ``Parameter.coerce_value``; - check other JSON values against the declared type; pass live Python objects - through unchanged. + resolve registry-reference parameters by name; build structured inputs from + their declared variants; coerce simple string values via + ``Parameter.coerce_value``; convert other JSON values via + ``Parameter.coerce_json_value``, which rejects JSON for parameters that take + only Python objects; pass live Python objects through unchanged. Args: cls (type): The class being built. @@ -592,7 +570,7 @@ def resolve_constructor_args( Raises: ValueError: If an argument is not a declared parameter, a registry reference cannot be resolved, a simple value cannot be coerced, or a - JSON value does not match the declared type. + JSON value is not accepted for the parameter. """ by_name = {param.name: param for param in derive_parameters(cls=cls, identifier_type=identifier_type)} @@ -604,7 +582,6 @@ def resolve_constructor_args( f"Unknown parameter '{name}' for '{cls.__name__}'. Valid parameters: {sorted(by_name.keys())}" ) - value_type = _unwrap_optional(param.param_type) if param.reference is not None: getter = _registry_getter_for_component_type(param.reference.component_type) if getter is None: @@ -621,51 +598,25 @@ def resolve_constructor_args( ) elif param.variants is not None: resolved[name] = _resolve_structured_input(parameter=param, value=value) - elif ( - (isinstance(value, str) and param.is_string_coercible) - or (isinstance(value_type, type) and issubclass(value_type, Enum)) - or (_is_choice_list(value_type) and _is_json_data(value)) - ): + elif isinstance(value, str) and param.is_string_coercible: try: resolved[name] = param.coerce_value(value) except (ValueError, TypeError) as e: raise ValueError(f"Parameter '{name}' of '{cls.__name__}': {e}") from e + elif _is_json_data(value): + resolved[name] = param.coerce_json_value(value, owner=cls.__name__) else: - resolved[name] = _check_json_value(parameter=param, value=value, owner=cls.__name__) + resolved[name] = value return resolved -def _is_choice_list(annotation: Any) -> bool: - """ - Return whether an annotation is a list of enum or literal choices. - - ``Parameter.coerce_value`` converts such a list from its JSON choice values, as it - does a single enum. - - Returns: - bool: True for ``list[SomeEnum]`` and ``list[Literal[...]]``. - """ - if get_origin(annotation) is not list: - return False - args = get_args(annotation) - element = args[0] if args else None - return get_origin(element) is Literal or (isinstance(element, type) and issubclass(element, Enum)) - - -_PLAIN_SEQUENCE_TYPES = (list, tuple, set, frozenset) - - -def _is_json_data(value: Any, *, sequence_types: tuple[type, ...] = (list,)) -> bool: +def _is_json_data(value: Any) -> bool: """ - Return whether a value holds only JSON-style data rather than live Python objects. - - Args: - value (Any): The value to inspect. - sequence_types (tuple[type, ...]): Sequence types accepted alongside string-keyed dicts. + Return whether a value holds only JSON data rather than live Python objects. Returns: - bool: True for None, strings, numbers, and the given sequences or string-keyed dicts of + bool: True for None, strings, numbers, booleans, and lists or string-keyed dicts of them (exact built-in types, so subclasses and enums count as live objects), without cycles. """ @@ -675,7 +626,7 @@ def _is_json_data(value: Any, *, sequence_types: tuple[type, ...] = (list,)) -> item = pending.pop() if item is None or type(item) in (str, int, float, bool): continue - if id(item) in seen or type(item) not in (dict, *sequence_types): + if id(item) in seen or type(item) not in (dict, list): return False seen.add(id(item)) if type(item) is dict: @@ -687,109 +638,6 @@ def _is_json_data(value: Any, *, sequence_types: tuple[type, ...] = (list,)) -> return True -def _json_projection(annotation: Any) -> Any: - """ - Project an annotation onto the types its JSON form is validated against. - - Abstract collections become lists and regex patterns become strings at any depth, - including inside unions and containers, so pydantic can validate their JSON form. - Protocols become a check that the value provides every protocol member. - - Returns: - Any: The projected annotation. - """ - origin = get_origin(annotation) or annotation - if origin is re.Pattern: - return str - if is_protocol(origin): - return Annotated[object, AfterValidator(functools.partial(_require_protocol_members, protocol=origin))] - args = get_args(annotation) - if origin in (Collection, Sequence, Iterable): - if not args: - return list - element_type = _json_projection(args[0]) - return list[element_type] - if not args: - return annotation - projected = tuple(_json_projection(arg) for arg in args) - if origin in (Union, types.UnionType): - return functools.reduce(operator.or_, projected) - return origin[projected] - - -def _require_protocol_members(value: object, *, protocol: type) -> object: - """ - Require a value to provide every member of a protocol. - - Returns: - object: The unchanged value. - - Raises: - ValueError: If the value lacks a protocol member. - """ - missing = sorted(name for name in get_protocol_members(protocol) if not hasattr(value, name)) - if missing: - raise ValueError(f"{protocol.__name__} requires {', '.join(missing)}") - return value - - -def _json_adapter(annotation: Any) -> TypeAdapter[Any]: - """ - Build a JSON validator for a parameter annotation. - - Annotations naming arbitrary classes are retried as instance checks, which no JSON value - satisfies, so JSON input for an object-only parameter is rejected instead of skipped. - - Returns: - TypeAdapter[Any]: The validator. - """ - try: - return TypeAdapter(annotation) - except PydanticSchemaGenerationError: - return TypeAdapter(annotation, config=ConfigDict(arbitrary_types_allowed=True)) - - -def _check_json_value(*, parameter: Parameter, value: Any, owner: str) -> Any: - """ - Check a JSON value against the parameter's declared type before construction. - - JSON values must match the declared type, and must not need conversion into objects such - as paths, enums, or models, which the constructor would otherwise receive as raw JSON. - JSON arrays become tuples or sets where the type declares them, and protocol-typed values - must provide the protocol's members. Live Python objects from in-process callers, and - annotations pydantic cannot resolve or build (such as forward references), are left to the - constructor. - - Args: - parameter (Parameter): The declared parameter. - value (Any): The raw supplied value. - owner (str): The class name used in error messages. - - Returns: - Any: The value to pass to the constructor. - - Raises: - ValueError: If the JSON value does not match the declared type. - """ - if parameter.param_type is None or not _is_json_data(value): - return value - try: - adapter = _json_adapter(_json_projection(parameter.param_type)) - validated = adapter.validate_json(json.dumps(value), strict=True) - plain = _is_json_data(validated, sequence_types=_PLAIN_SEQUENCE_TYPES) - unchanged = plain and validated == value - except ValidationError as exc: - raise ValueError(f"Parameter '{parameter.name}' of '{owner}' expects {parameter.type_name}.") from exc - except (PydanticUserError, PydanticUndefinedAnnotation, SchemaError, RecursionError, TypeError, ValueError): - logger.debug("Skipping JSON type check for %s.%s", owner, parameter.name) - return value - if not plain: - raise ValueError( - f"Parameter '{parameter.name}' of '{owner}' expects {parameter.type_name}, which cannot be built from JSON." - ) - return value if unchanged else validated - - def _resolve_structured_input(*, parameter: Parameter, value: Any) -> Any: """ Build a declared structured-input variant from its JSON representation. diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index 7d34ecc42a..4940e42da7 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -213,19 +213,21 @@ async def test_list_converter_types_includes_supported_types(self) -> None: assert "text" in base64_entry.supported_input_types assert "text" in base64_entry.supported_output_types - async def test_types_include_all_constructible_converters(self) -> None: - """The projection surfaces every constructible converter, including base/helper classes. + async def test_types_include_every_registered_converter(self) -> None: + """The projection surfaces every registered converter, including base/helper classes. Whether to display a given converter is left to the caller (e.g. the frontend), - so the service no longer hides anything. + so the service hides nothing and reports which ones the API can construct. """ service = ConverterService() result = await service.list_converter_types_async() - converter_types = [item.converter_type for item in result.items] - assert "Base64Converter" in converter_types - assert "SelectiveTextConverter" in converter_types + constructible = {item.converter_type: item.constructible for item in result.items} + assert constructible["Base64Converter"] is True + assert constructible["SearchReplaceConverter"] is True + assert constructible["SelectiveTextConverter"] is False + assert constructible["TextJailbreakConverter"] is False async def test_types_serialize_parameter_type(self) -> None: """Type entries render the raw annotation into a human-readable type_name.""" diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index ac997198b2..5477ac8e1f 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -247,6 +247,16 @@ async def test_types_return_known_target_types(self) -> None: assert "OpenAIChatTarget" in target_types assert "AzureMLChatTarget" in target_types + async def test_types_report_whether_required_parameters_can_be_supplied(self) -> None: + service = TargetService() + + result = await service.list_target_types_async() + + constructible = {item.target_type: item.constructible for item in result.items} + assert constructible["OpenAIChatTarget"] is True + assert constructible["WebsocketTarget"] is False + assert constructible["PlaywrightTarget"] is False + async def test_types_include_declarative_auth_facts(self) -> None: """Type entries surface the per-class auth facts the frontend needs.""" service = TargetService() diff --git a/tests/unit/models/test_parameter.py b/tests/unit/models/test_parameter.py index 6e2e3d0785..cf7e8421c8 100644 --- a/tests/unit/models/test_parameter.py +++ b/tests/unit/models/test_parameter.py @@ -3,9 +3,10 @@ """Unit tests for the unified Parameter model and its coercion methods.""" +from collections.abc import Callable, Mapping, Sequence from enum import Enum from pathlib import Path -from typing import Literal, Union +from typing import Any, Literal, Protocol, Union import pytest from pydantic import ValidationError @@ -86,6 +87,7 @@ def test_scalar_with_default(self) -> None: "choices": None, "is_list": False, "reference_type": None, + "input_kind": "scalar", "variants": None, } @@ -521,3 +523,155 @@ def __init__(self, *, value=None) -> None: param = next(p for p in derive_parameters(cls=_Holder) if p.name == "value") assert param.coerce_value(raw) == expected + + +class _Greeter(Protocol): + def greet(self) -> str: ... + + +class TestInputKind: + """``input_kind`` states how registry callers supply each parameter.""" + + @pytest.mark.parametrize( + ("param_type", "expected"), + [ + (str, "scalar"), + (int | None, "scalar"), + (Path, "scalar"), + (Path | str | None, "scalar"), + (Literal["a", "b"], "scalar"), + (_Speed, "scalar"), + (_Speed | Path, "scalar"), + (_Unsupported | str, "scalar"), + (list[str], "collection"), + (tuple[int, int], "collection"), + (set[str], "collection"), + (Sequence[str] | None, "collection"), + (dict[str, Any], "collection"), + (Mapping[str, list[str]], "collection"), + (list[dict[str, Any]], "collection"), + (str | list[str], "collection"), + (_Unsupported, "in_process_only"), + (Callable[[str], str], "in_process_only"), + (_Greeter | None, "in_process_only"), + (list[Path], "in_process_only"), + (Sequence[Path | str], "in_process_only"), + (dict[int, str], "in_process_only"), + (list[str | bytes], "in_process_only"), + (None, "unsupported"), + (Any, "unsupported"), + ("SeedPrompt | None", "unsupported"), + (list["SeedPrompt"], "unsupported"), + ], + ) + def test_value_parameter_kinds(self, param_type: object, expected: str) -> None: + parameter = Parameter(name="p", description="d", param_type=param_type) + + assert parameter.input_kind == expected + assert parameter.model_dump()["input_kind"] == expected + assert parameter.is_json_configurable is (expected in ("scalar", "collection")) + + @pytest.mark.parametrize( + ("param_type", "expected"), + [ + (int, "scalar"), + (list[str], "collection"), + (tuple[int, int], "collection"), + (_Unsupported, "in_process_only"), + ("SeedPrompt | None", "unsupported"), + ], + ) + def test_input_kind_survives_wire_round_trip(self, param_type: object, expected: str) -> None: + restored = Parameter.model_validate(Parameter(name="p", description="d", param_type=param_type).model_dump()) + + assert restored.input_kind == expected + assert restored.model_dump()["input_kind"] == expected + + def test_reference_structured_and_opaque_kinds(self) -> None: + reference = Parameter( + name="t", description="d", reference=RegistryReference(component_type=ComponentType.TARGET) + ) + structured = Parameter(name="s", description="d", param_type=_Unsupported, variants={"one": []}) + opaque = Parameter(name="o", description="d", param_type=str, opaque=True) + + assert (reference.input_kind, structured.input_kind, opaque.input_kind) == ( + "reference", + "structured", + "in_process_only", + ) + assert reference.is_json_configurable and structured.is_json_configurable + assert not opaque.is_json_configurable + + +class TestCoerceJsonValue: + """``coerce_json_value`` applies the registry contract to JSON input.""" + + @pytest.mark.parametrize( + ("param_type", "value", "expected"), + [ + (int, 3, 3), + (float, 2, 2.0), + (str | None, None, None), + (_Speed, "fast", _Speed.FAST), + (Literal[1, 2], "2", 2), + (Path, "/data/input.png", Path("/data/input.png")), + (Path | str, "relative.png", "relative.png"), + (tuple[int, int], [1, 2], (1, 2)), + (tuple[str, ...], ["a", "b"], ("a", "b")), + (set[str], ["a", "a"], {"a"}), + (frozenset[int], [1], frozenset({1})), + (list[_Speed], ["slow"], [_Speed.SLOW]), + (dict[str, list[str]], {"k": ["v"]}, {"k": ["v"]}), + (str | list[str], "x", "x"), + (str | list[str], ["x"], ["x"]), + (_Speed | Path, "slow", _Speed.SLOW), + (_Unsupported | str, "text", "text"), + ], + ) + def test_converts_matching_json(self, param_type: object, value: object, expected: object) -> None: + parameter = Parameter(name="p", description="d", param_type=param_type) + + converted = parameter.coerce_json_value(value, owner="Owner") + + assert converted == expected + assert type(converted) is type(expected) + + @pytest.mark.parametrize( + ("param_type", "value"), + [ + (int, True), + (int, 1.5), + (bool, 1), + (str, 5), + (float, "1.5"), + (int, None), + (tuple[int, int], [1]), + (list[str], "abc"), + (list[Path], ["/etc/hostname"]), + (_Speed | Path, "/etc/hostname"), + (dict[str, int], {"k": "v"}), + (set[str], [["nested"]]), + (list[_Speed], ["bogus"]), + (float, 10**400), + ], + ) + def test_rejects_mismatched_json(self, param_type: object, value: object) -> None: + parameter = Parameter(name="p", description="d", default=REQUIRED_VALUE, param_type=param_type) + + with pytest.raises(ValueError, match="Parameter 'p' of 'Owner'"): + parameter.coerce_json_value(value, owner="Owner") + + @pytest.mark.parametrize( + ("param_type", "message"), + [(_Unsupported, "accepts only a Python object"), (Any, "does not support"), (None, "does not support")], + ) + def test_rejects_json_for_parameters_without_json_input(self, param_type: object, message: str) -> None: + parameter = Parameter(name="p", description="d", param_type=param_type) + + with pytest.raises(ValueError, match=message): + parameter.coerce_json_value({"a": 1}, owner="Owner") + + def test_none_is_accepted_when_default_is_none(self) -> None: + parameter = Parameter(name="p", description="d", default=None, param_type=_Unsupported) + + assert parameter.coerce_json_value(None, owner="Owner") is None diff --git a/tests/unit/registry/test_registry_metadata.py b/tests/unit/registry/test_registry_metadata.py index 40007babfe..b6152b7890 100644 --- a/tests/unit/registry/test_registry_metadata.py +++ b/tests/unit/registry/test_registry_metadata.py @@ -3,6 +3,8 @@ from dataclasses import dataclass, field +from pyrit.common import REQUIRED_VALUE +from pyrit.models import Parameter from pyrit.registry.registry import _matches_filters from pyrit.registry.registry_metadata import RegistryMetadata @@ -223,3 +225,23 @@ def test_matches_filters_combined_include_and_exclude(self): ) is False ) + + +class TestConstructible: + """``constructible`` reports whether every required parameter takes JSON input.""" + + def test_optional_object_parameter_keeps_class_constructible(self) -> None: + parameters = ( + Parameter(name="count", description="", default=REQUIRED_VALUE, param_type=int), + Parameter(name="handle", description="", default=None, param_type=object), + ) + + assert RegistryMetadata(class_name="C", class_module="m", parameters=parameters).constructible + + def test_required_object_parameter_makes_class_not_constructible(self) -> None: + parameters = ( + Parameter(name="count", description="", default=REQUIRED_VALUE, param_type=int), + Parameter(name="handle", description="", default=REQUIRED_VALUE, param_type=object), + ) + + assert not RegistryMetadata(class_name="C", class_module="m", parameters=parameters).constructible diff --git a/tests/unit/registry/test_resolution.py b/tests/unit/registry/test_resolution.py index dd70877173..6caa1a4112 100644 --- a/tests/unit/registry/test_resolution.py +++ b/tests/unit/registry/test_resolution.py @@ -350,6 +350,7 @@ def test_unknown_registry_reference_empty_registry_hint(self, empty_target_regis (_SimpleOnly, {"count": None}), (_SimpleOnly, {"ratio": {"value": 1}}), (_SimpleOnly, {"flag": 1}), + (_SimpleOnly, {"ratio": 10**400}), (_JsonShaped, {"note": 5}), (_JsonShaped, {"weights": [1, "2"]}), (_JsonShaped, {"extra": "not-an-object"}), @@ -388,10 +389,16 @@ def test_accepts_matching_json_value_unchanged(self, cls: type, raw_args: dict[s def test_json_array_becomes_tuple_for_tuple_parameter(self) -> None: assert _resolve(_JsonShaped, {"color": [10, 20, 30]})["color"] == (10, 20, 30) - @pytest.mark.parametrize("raw_args", [{"settings": {"level": 1}}, {"location": "fast"}, {"location": "/tmp/x"}]) - def test_rejects_json_value_that_needs_an_object(self, raw_args: dict[str, object]) -> None: - with pytest.raises(ValueError, match="cannot be built from JSON"): - _resolve(_JsonShaped, raw_args) + def test_rejects_json_for_in_process_only_parameter(self) -> None: + with pytest.raises(ValueError, match="settings.*accepts only a Python object"): + _resolve(_JsonShaped, {"settings": {"level": 1}}) + + def test_union_takes_json_only_through_its_json_members(self) -> None: + assert _resolve(_JsonShaped, {"location": "fast"}) == {"location": _Speed.FAST} + with pytest.raises(ValueError, match="location"): + _resolve(_JsonShaped, {"location": "/tmp/x"}) + path = Path("/tmp/x") + assert _resolve(_JsonShaped, {"location": path})["location"] is path def test_live_objects_pass_through_unchecked(self) -> None: color = object() @@ -414,16 +421,23 @@ def test_deeply_nested_json_value_is_rejected(self) -> None: with pytest.raises(ValueError, match="expects"): _resolve(_SimpleOnly, {"count": nested}) - def test_protocol_parameter_requires_protocol_members(self) -> None: - with pytest.raises(ValueError, match="provider"): - _resolve(_JsonShaped, {"provider": "anything"}) - with pytest.raises(ValueError, match="sized"): - _resolve(_JsonShaped, {"sized": 7}) + @pytest.mark.parametrize("raw_args", [{"provider": "anything"}, {"sized": 7}, {"sized": [1, 2]}]) + def test_protocol_parameter_takes_only_live_objects(self, raw_args: dict[str, object]) -> None: + with pytest.raises(ValueError, match="accepts only a Python object"): + _resolve(_JsonShaped, raw_args) + + def test_protocol_parameter_accepts_live_object_and_none(self) -> None: + sized = _Bag([1, 2]) + + assert _resolve(_JsonShaped, {"provider": None, "sized": sized}) == {"provider": None, "sized": sized} - assert _resolve(_JsonShaped, {"provider": None, "sized": [1, 2]}) == {"provider": None, "sized": [1, 2]} + def test_unresolved_annotation_rejects_json(self) -> None: + target = MockPromptTarget() - def test_unresolved_annotation_is_left_to_constructor(self) -> None: - assert _resolve(_Unresolved, {"target": {"a": 1}}) == {"target": {"a": 1}} + with pytest.raises(ValueError, match="does not support"): + _resolve(_Unresolved, {"target": {"a": 1}}) + assert _resolve(_Unresolved, {"target": None}) == {"target": None} + assert _resolve(_Unresolved, {"target": target})["target"] is target def test_json_value_for_object_parameter_is_rejected(self) -> None: handle = _Handle() @@ -434,20 +448,38 @@ def test_json_value_for_object_parameter_is_rejected(self) -> None: assert _resolve(_JsonShaped, {"handle": handle})["handle"] is handle @pytest.mark.parametrize( - ("registry_type", "identifier_type", "type_name", "raw_args"), + ("registry_type", "identifier_type", "type_name", "raw_args", "message"), [ - (TargetRegistry, TargetIdentifier, "TextTarget", {"custom_configuration": {}}), - (TargetRegistry, TargetIdentifier, "A2ATarget", {"auth_token": {"token": "x"}}), - (TargetRegistry, TargetIdentifier, "OpenAIResponseTarget", {"tool_providers": [{"name": "x"}]}), - (ConverterRegistry, ConverterIdentifier, "TokenBijectionConverter", {"tokenizer": "name"}), + (TargetRegistry, TargetIdentifier, "TextTarget", {"custom_configuration": {}}, "accepts only"), + (TargetRegistry, TargetIdentifier, "A2ATarget", {"auth_token": {"token": "x"}}, "expects"), + ( + TargetRegistry, + TargetIdentifier, + "OpenAIResponseTarget", + {"tool_providers": [{"name": "x"}]}, + "accepts only", + ), + (ConverterRegistry, ConverterIdentifier, "TokenBijectionConverter", {"tokenizer": "name"}, "accepts only"), + ( + ConverterRegistry, + ConverterIdentifier, + "TextJailbreakConverter", + {"jailbreak_template": {}}, + "accepts only", + ), ], ) def test_registered_component_rejects_json_for_object_parameter( - self, registry_type: type, identifier_type: type, type_name: str, raw_args: dict[str, object] + self, + registry_type: type, + identifier_type: type, + type_name: str, + raw_args: dict[str, object], + message: str, ) -> None: cls = registry_type.get_registry_singleton().get_class(type_name) - with pytest.raises(ValueError, match="expects"): + with pytest.raises(ValueError, match=message): _resolve(cls, raw_args, identifier_type=identifier_type) @pytest.mark.parametrize( @@ -506,7 +538,7 @@ def test_registered_json_defaults_are_accepted(self, registry_type: type, identi for type_name in registry.get_class_names(): cls = registry.get_class(type_name) for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): - if parameter.reference is not None or parameter.default is None: + if parameter.input_kind not in ("scalar", "collection") or parameter.default is None: continue try: value = json.loads(json.dumps(parameter.default)) From c1431885a445bf554443b900eaf6a6cbb72827f7 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Mon, 5 Oct 2026 17:24:48 +0000 Subject: [PATCH 3/8] Narrow registry validation to the supported external-input boundary Component validation stays in constructors. The registry only enforces which inputs external callers may supply. Removes the input_kind classification, the recursive JSON conversion, the value-shape caller detection, and the constructible and wire metadata. External callers now go through an explicit path (Registry.create_instance_from_external_input, or create_named_instance with external_input=True). It accepts only parameters with Parameter.is_external_input: scalars, lists of non-path scalars, unions of str with those or with callables (such as api_key: str | Callable), registry references given by name, and declared structured inputs. Unknown names are rejected. In-process callers can still pass any Python object. The converter, target, and scenario catalogs list only those inputs. Scenario runs and estimates apply the same check to scenario_params before the server merges its own values. framework.md states the rule in one sentence. --- doc/code/framework.md | 11 +- pyrit/backend/models/converters.py | 3 - pyrit/backend/models/targets.py | 3 - pyrit/backend/services/converter_service.py | 14 +- .../backend/services/scenario_run_service.py | 9 +- pyrit/backend/services/scenario_service.py | 9 +- pyrit/backend/services/target_service.py | 13 +- pyrit/converter/word_doc_converter.py | 3 +- pyrit/models/parameter.py | 358 ++-------------- pyrit/registry/registry.py | 42 +- pyrit/registry/registry_metadata.py | 5 - pyrit/registry/resolution.py | 118 +++--- tests/unit/backend/test_converter_service.py | 60 ++- .../unit/backend/test_scenario_run_service.py | 33 ++ tests/unit/backend/test_scenario_service.py | 49 +++ tests/unit/backend/test_target_service.py | 58 ++- .../unit/converter/test_word_doc_converter.py | 8 +- tests/unit/models/test_parameter.py | 180 ++------ .../unit/registry/test_converter_registry.py | 18 + tests/unit/registry/test_registry_metadata.py | 22 - tests/unit/registry/test_resolution.py | 398 +++++------------- .../score/test_garak_exploitation_scorer.py | 7 +- 22 files changed, 508 insertions(+), 913 deletions(-) diff --git a/doc/code/framework.md b/doc/code/framework.md index a0c26a8a23..b03dd4b802 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -394,16 +394,7 @@ See [message normalizers](./targets/11_message_normalizer) for capability behavi - If you are creating a component with user input (e.g. via config, REST, or automatically) it should always use the registry - If you are storing an instance of a component, it should always use the registry -- Constructor parameters are the registry's input contract. Each derived `Parameter` has an `input_kind` that says how callers supply it: - - `scalar`: `str`, `int`, `float`, `bool`, `Path`, `Path | str`, `Literal[...]`, or an `Enum`, optionally nullable. JSON callers send one string, number, or boolean; choices match by value or name. - - `collection`: a list, tuple, set, `Sequence`, `Collection`, or string-keyed `dict` / `Mapping` of scalar or collection values, never `Path`. JSON callers send an array or object; arrays become the declared tuple or set. - - `reference`: a parameter the component's identifier marks as a registry reference. Callers send a registry name, or a list of names. - - `structured`: a `StructuredParameterValue`. Callers send one of its declared variants. - - `in_process_only`: any other type (a class, callable, or protocol), or a parameter marked `opaque`. JSON is rejected; Python callers pass a live object. - - `unsupported`: no annotation, `Any`, or an annotation that cannot be resolved at runtime. JSON is rejected; resolve the annotation to make the parameter configurable. -- A union takes JSON when one of its members does (e.g. `SeedPrompt | str` takes a string), and a JSON value becomes the first member it matches. A `Path` inside a wider union or a collection is not JSON input, so JSON cannot name server files. -- The resolver enforces this contract for every caller: JSON for an `in_process_only` or `unsupported` parameter is rejected, not passed to the constructor. Live Python objects pass through unchanged, and constructors still own component-specific validation. -- A component with a required parameter that is not JSON input is not `constructible`; registry metadata and the REST catalog report this so clients do not offer it. +- The registry accepts only explicitly supported external inputs, permits opaque Python objects only for in-process callers, and leaves component validation to constructors. ## [Setup](./setup/0_setup) diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 057a3ff2fa..316cedb42b 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -44,9 +44,6 @@ class ConverterTypeEntry(BaseModel): parameters: list[Parameter] = Field( default_factory=list, description="Constructor parameters for dynamic form generation" ) - constructible: bool = Field( - True, description="Whether every required parameter can be supplied through this API (see input_kind)" - ) is_llm_based: bool = Field(False, description="Whether this converter requires an LLM target") description: str | None = Field(None, description="Short description of the converter from its docstring") diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index 9479d042b9..ba9759e767 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -36,9 +36,6 @@ class TargetTypeEntry(BaseModel): default_factory=list, description="Constructor parameters for dynamic form generation", ) - constructible: bool = Field( - True, description="Whether every required parameter can be supplied through this API (see input_kind)" - ) supported_auth_modes: list[Literal["api_key", "identity"]] = Field( default_factory=_default_auth_modes, description="Authentication modes this target type supports", diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index 3e060bec67..165a0a4efc 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -112,10 +112,11 @@ async def list_converter_types_async(self) -> ConverterTypeResponse: """ List all available converter types from the converter class registry. - Returns every registered converter type. ``constructible`` is False when a - required parameter cannot be supplied through the API. Deciding which - entries to surface to a user is a presentation concern owned by the caller - (e.g. the frontend), not this service. + Returns every converter that external callers can build, with only the + parameters they may supply; converters that need a Python object for a + required parameter are left out. Deciding which entries to surface to a + user is a presentation concern owned by the caller (e.g. the frontend), + not this service. Returns: ConverterTypeResponse containing all available converter classes. @@ -125,12 +126,12 @@ async def list_converter_types_async(self) -> ConverterTypeResponse: converter_type=metadata.class_name, supported_input_types=list(metadata.supported_input_types), supported_output_types=list(metadata.supported_output_types), - parameters=list(metadata.parameters), - constructible=metadata.constructible, + parameters=[parameter for parameter in metadata.parameters if parameter.is_external_input], is_llm_based=metadata.is_llm_based, description=metadata.class_description or None, ) for metadata in self._registry.get_all_registered_class_metadata() + if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) ] return ConverterTypeResponse(items=items) @@ -201,6 +202,7 @@ async def create_converter_async(self, *, request: CreateConverterRequest) -> Co type_name=request.type, params=params, registry_metadata={_OWNED_ARTIFACT_PATHS_KEY: [str(path) for path in owned_paths]}, + external_input=True, ) except (Exception, asyncio.CancelledError): await self._remove_owned_artifacts_async(paths=owned_paths) diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index 50808b204c..b7838a4b0f 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -78,7 +78,7 @@ ) from pyrit.prompt_target import PromptTarget from pyrit.registry import InitializerRegistry, ScenarioRegistry -from pyrit.registry.resolution import resolve_declared_params +from pyrit.registry.resolution import reject_non_external_params, resolve_declared_params from pyrit.scenario import Scenario from pyrit.scenario.core import override_default_adversarial_target @@ -246,6 +246,13 @@ async def start_run_async(self, *, request: RunScenarioRequest) -> ScenarioRunSu Returns: ScenarioRunSummary: Current scheduled run state. """ + registry = ScenarioRegistry.get_registry_singleton() + if request.scenario_params and request.scenario_name in registry: + reject_non_external_params( + params=request.scenario_params, + declared=registry.get_class(request.scenario_name).supported_parameters(), + owner=request.scenario_name, + ) async with self._reserve_resume_request_async(request.scenario_result_id), self._launch_lock: await self._validate_resume_admission_async(scenario_result_id=request.scenario_result_id) return await self._start_run_locked_async(request=request) diff --git a/pyrit/backend/services/scenario_service.py b/pyrit/backend/services/scenario_service.py index 1596e2c59b..4d18a4eb68 100644 --- a/pyrit/backend/services/scenario_service.py +++ b/pyrit/backend/services/scenario_service.py @@ -21,6 +21,7 @@ ScenarioRunSizeEstimateRequest, ) from pyrit.registry import ScenarioMetadata, ScenarioRegistry +from pyrit.registry.resolution import reject_non_external_params from pyrit.scenario.core import Scenario, override_default_adversarial_target from pyrit.scenario.core.dataset_configuration import read_only_dataset_resolution @@ -68,7 +69,7 @@ def _metadata_to_registered_scenario( all_techniques=list(metadata.all_techniques), technique_summaries=list(metadata.technique_summaries), default_datasets=list(metadata.default_datasets), - supported_parameters=list(metadata.supported_parameters), + supported_parameters=[parameter for parameter in metadata.supported_parameters if parameter.is_external_input], baseline_policy=metadata.baseline_policy, include_baseline_by_default=metadata.include_baseline_by_default, uses_default_adversarial_target=metadata.uses_default_adversarial_target, @@ -193,6 +194,12 @@ async def estimate_scenario_run_size_async( scenario_class = self._registry.get_class(scenario_name) except KeyError: return None + if request.scenario_params: + reject_non_external_params( + params=request.scenario_params, + declared=scenario_class.supported_parameters(), + owner=scenario_name, + ) estimate_key = self._build_configured_estimate_key( scenario_name=scenario_name, scenario_class=scenario_class, request=request diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 8942cdd8f3..f8c20bd2e5 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -152,10 +152,10 @@ async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. - Returns every registered target type with its derived constructor - parameters, whether its required parameters can be supplied through the - API (``constructible``), and the auth modes it supports, all projected from - the registry's ``TargetMetadata``. Deciding which entries to surface to a + Returns every target that external callers can build, with the + constructor parameters they may supply and the auth modes it supports, + all projected from the registry's ``TargetMetadata``; targets that need a + Python object for a required parameter are left out. Deciding which entries to surface to a user is a presentation concern owned by the caller (e.g. the frontend), not this service. @@ -166,12 +166,12 @@ async def list_target_types_async(self) -> TargetTypeResponse: items: list[TargetTypeEntry] = [ TargetTypeEntry( target_type=metadata.class_name, - parameters=list(metadata.parameters), - constructible=metadata.constructible, + parameters=[parameter for parameter in metadata.parameters if parameter.is_external_input], supported_auth_modes=self._get_supported_auth_modes(metadata.supported_auth_modes), description=metadata.class_description or None, ) for metadata in metadata_items + if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) ] return TargetTypeResponse(items=items) @@ -220,6 +220,7 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn name=target_registry_name, type_name=request.type, params=params, + external_input=True, ) return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) diff --git a/pyrit/converter/word_doc_converter.py b/pyrit/converter/word_doc_converter.py index 02e109a6e2..a87c0cac39 100644 --- a/pyrit/converter/word_doc_converter.py +++ b/pyrit/converter/word_doc_converter.py @@ -15,11 +15,10 @@ from pyrit.common.logger import logger from pyrit.converter.converter import Converter, ConverterResult from pyrit.memory import data_serializer_factory -from pyrit.models import SeedPrompt # noqa: TC001 - registry annotation resolution if TYPE_CHECKING: from pyrit.memory import DataTypeSerializer - from pyrit.models import ComponentIdentifier, PromptDataType + from pyrit.models import ComponentIdentifier, PromptDataType, SeedPrompt @dataclass diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index edf170f02e..5370fabf6e 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -8,36 +8,16 @@ import copy import types from abc import ABC, abstractmethod -from collections.abc import Collection, Iterable, Mapping, MutableMapping, MutableSequence, MutableSet, Sequence -from collections.abc import Set as AbstractSet +from collections.abc import Callable from dataclasses import dataclass from enum import Enum from pathlib import Path -from typing import Any, ForwardRef, Literal, TypeAlias, Union, get_args, get_origin +from typing import Any, Literal, Union, get_args, get_origin from pydantic import BaseModel, ConfigDict, Field, computed_field, field_serializer, model_validator from pyrit.common.apply_defaults import REQUIRED_VALUE -#: How registry callers supply a parameter; see ``Parameter.input_kind``. -ParameterInputKind: TypeAlias = Literal[ - "scalar", "collection", "reference", "structured", "in_process_only", "unsupported" -] -_JSON_INPUT_KINDS: frozenset[str] = frozenset({"scalar", "collection", "reference", "structured"}) -_SEQUENCE_ORIGINS: tuple[Any, ...] = ( - list, - tuple, - set, - frozenset, - Collection, - Sequence, - MutableSequence, - Iterable, - AbstractSet, - MutableSet, -) -_SET_ORIGINS: tuple[Any, ...] = (set, frozenset, AbstractSet, MutableSet) -_MAPPING_ORIGINS: tuple[Any, ...] = (dict, Mapping, MutableMapping) _SUPPORTED_SCALAR_TYPES: tuple[type, ...] = (str, int, float, bool, Path) _SCALAR_NAME_TO_TYPE: dict[str, type | types.UnionType] = { "Path": Path, @@ -106,8 +86,7 @@ class Parameter(BaseModel): *is* the allowed set) and drives ``coerce_value`` / ``validate``; it is **not** serialized. Serialization instead projects the type into the display fields ``type_name``, ``choices``, and ``is_list`` (plus ``required`` from the - ``REQUIRED_VALUE`` sentinel and ``input_kind``, which says how registry callers - supply the value), so a consumer can rebuild a usable contract from + ``REQUIRED_VALUE`` sentinel), so a consumer can rebuild a usable contract from the registry without the live type travelling on the wire. ``reference``, when set, marks the parameter as a registry reference: its value @@ -151,14 +130,6 @@ class Parameter(BaseModel): exclude=True, description="Where the parameter is consumed at build time; not serialized.", ) - wire_input_kind: ParameterInputKind | None = Field( - default=None, - exclude=True, - description=( - "The ``input_kind`` read from a serialized payload. A ``param_type`` rebuilt from display " - "fields cannot always express it, so the serialized value is kept. Not serialized." - ), - ) opaque: bool = Field( default=False, exclude=True, @@ -201,7 +172,6 @@ def _reconstruct_param_type_from_wire(cls, data: Any) -> Any: choices=data.get("choices"), is_list=bool(data.get("is_list")), ) - data.setdefault("wire_input_kind", data.get("input_kind")) if needs_reference: data["reference"] = RegistryReference( component_type=ComponentType(data["reference_type"]), @@ -258,33 +228,6 @@ def reference_type(self) -> str | None: """Registry component family this parameter references, or None.""" return self.reference.component_type.value if self.reference is not None else None - @computed_field - @property - def input_kind(self) -> ParameterInputKind: - """ - How registry callers supply this parameter. - - ``scalar`` and ``collection`` parameters take JSON values of their declared type, - ``reference`` parameters take registry names, and ``structured`` parameters take one of - their declared variants. ``in_process_only`` parameters take only live Python objects, - and ``unsupported`` parameters have a type the registry cannot describe; neither takes - JSON. - """ - if self.wire_input_kind is not None: - return self.wire_input_kind - if self.reference is not None: - return "reference" - if self.variants is not None: - return "structured" - if self.opaque: - return "in_process_only" - return _annotation_input_kind(self.param_type) - - @property - def is_json_configurable(self) -> bool: - """Whether registry callers can supply this parameter as JSON.""" - return self.input_kind in _JSON_INPUT_KINDS - @field_serializer("default") def _serialize_default(self, value: Any) -> str | list[str] | None: """ @@ -322,6 +265,34 @@ def is_string_coercible(self) -> bool: return False return _is_scalar_param_type(_unwrap_optional(self.param_type)) + @property + def is_external_input(self) -> bool: + """ + Whether REST, CLI, and GUI callers may supply this parameter. + + True for registry references, declared structured inputs, scalars, lists of + non-path scalars, and unions of ``str`` with those or with callables (external + callers send the string, as for ``api_key: str | Callable[...]``; callables are + for in-process callers). Other parameters take Python objects from in-process + callers only. + + Returns: + bool: True when external callers may supply this parameter. + """ + if self.reference is not None or self.variants is not None: + return True + if self.opaque: + return False + param_type = _unwrap_optional(self.param_type) + if _is_scalar_param_type(param_type): + return True + if get_origin(param_type) in (Union, types.UnionType): + members = [member for member in get_args(param_type) if member is not type(None)] + return str in members and all( + _is_non_path_json_type(member) or get_origin(member) is Callable for member in members + ) + return _is_non_path_json_type(param_type) + def is_reference_to(self, component_type: ComponentType) -> bool: """ Whether this parameter is a registry reference to the given component family. @@ -376,43 +347,6 @@ def coerce_value(self, raw_value: Any) -> Any: return _coerce_simple_value(param_name=self.name, annotation=param_type, raw_value=raw_value) return raw_value - def coerce_json_value(self, value: Any, *, owner: str) -> Any: - """ - Convert a JSON value to this parameter's declared type under the registry contract. - - A ``scalar`` or ``collection`` parameter takes JSON of its declared type: arrays become - the declared list, tuple, or set, and enum or literal choices become their members. A - path is accepted only where the parameter is itself a ``Path`` or ``Path | str``. Other - parameters take no JSON. ``None`` is accepted wherever the parameter allows it. - - Args: - value (Any): JSON data: None, a string, number, boolean, list, or string-keyed dict. - owner (str): The class name used in error messages. - - Returns: - Any: The value converted to the declared type. - - Raises: - ValueError: If the parameter cannot be supplied as JSON, or the value does not match - its declared type. - """ - if value is None and (self.default is None or type(None) in _union_members(self.param_type)): - return None - input_kind = self.input_kind - if input_kind == "in_process_only": - raise ValueError( - f"Parameter '{self.name}' of '{owner}' accepts only a Python object ({self.type_name}), not JSON." - ) - if input_kind not in ("scalar", "collection"): - raise ValueError( - f"Parameter '{self.name}' of '{owner}' has a type the registry does not support " - f"({self.type_name}), so it cannot be set from JSON." - ) - try: - return _convert_json_value(self.param_type, value, allow_path=True) - except _JsonMismatchError: - raise ValueError(f"Parameter '{self.name}' of '{owner}' expects {self.type_name}.") from None - def validate(self) -> None: # type: ignore[ty:invalid-method-override] """ Reject a declaration with an unsupported ``param_type``. @@ -493,233 +427,17 @@ def _is_scalar_param_type(annotation: Any) -> bool: return _is_enum_type(annotation) -class _JsonMismatchError(Exception): - """Raised when a JSON value does not match a declared type.""" - - -def _union_members(annotation: Any) -> tuple[Any, ...]: +def _is_non_path_json_type(annotation: Any) -> bool: """ - Return the members of a union annotation, or the annotation alone. + Return whether the annotation is a non-path scalar or a ``list`` of non-path scalars. Returns: - tuple[Any, ...]: The union's members, including ``NoneType`` when present. - """ - if get_origin(annotation) in (Union, types.UnionType): - return get_args(annotation) - return (annotation,) - - -def _is_unresolved(annotation: Any) -> bool: - """ - Return whether an annotation is, or contains, a forward reference that was never resolved. - - Returns: - bool: True when a string or ``ForwardRef`` stands in for a type. - """ - if isinstance(annotation, (str, ForwardRef)): - return True - if get_origin(annotation) is Literal: - return False - return any(_is_unresolved(arg) for arg in get_args(annotation)) - - -def _is_json_scalar_type(annotation: Any) -> bool: - """ - Return whether JSON expresses the annotation as one string, number, or boolean. - - Returns: - bool: True for ``str``, ``int``, ``float``, ``bool``, an ``Enum``, or a ``Literal`` of those values. - """ - if annotation in (str, int, float, bool) or _is_enum_type(annotation): - return True - return get_origin(annotation) is Literal and all( - arg is None or type(arg) in (str, int, float, bool) for arg in get_args(annotation) - ) - - -def _is_json_element_type(annotation: Any) -> bool: - """ - Return whether a collection element type can be expressed in JSON. - - Paths are not, so a collection cannot carry server file paths. - - Returns: - bool: True for JSON scalars, JSON collections, ``Any``, ``None``, and unions of those. - """ - if annotation is Any or annotation is type(None): - return True - if get_origin(annotation) in (Union, types.UnionType): - return all(_is_json_element_type(member) for member in get_args(annotation)) - return _is_json_scalar_type(annotation) or _is_json_collection_type(annotation) - - -def _is_json_collection_type(annotation: Any) -> bool: - """ - Return whether the annotation is a list, tuple, set, or string-keyed mapping of JSON values. - - Returns: - bool: True when a JSON array or object can express the annotation. - """ - origin = get_origin(annotation) or annotation - args = [arg for arg in get_args(annotation) if arg is not Ellipsis] - if origin in _MAPPING_ORIGINS: - key_type, value_type = args if len(args) == 2 else (str, Any) - return key_type in (str, Any) and _is_json_element_type(value_type) - if origin in _SEQUENCE_ORIGINS: - return all(_is_json_element_type(arg) for arg in args) - return False - - -def _member_input_kind(annotation: Any, *, in_union: bool) -> ParameterInputKind: - """ - Classify one member of a value parameter's annotation. - - A ``Path`` is a scalar only when it is the whole type; inside a wider union it can only be - passed as a live object. - - Returns: - ParameterInputKind: ``scalar``, ``collection``, ``in_process_only``, or ``unsupported``. - """ - if annotation is Any: - return "unsupported" - if annotation is Path: - return "in_process_only" if in_union else "scalar" - if _is_json_scalar_type(annotation): - return "scalar" - if _is_json_collection_type(annotation): - return "collection" - return "in_process_only" - - -def _annotation_input_kind(annotation: Any) -> ParameterInputKind: - """ - Classify a value parameter's annotation under the registry input contract. - - A union takes JSON when any member does; its other members can still be passed as live - objects. Unannotated parameters, ``Any``, and unresolved annotations are unsupported. - - Returns: - ParameterInputKind: ``scalar``, ``collection``, ``in_process_only``, or ``unsupported``. - """ - if annotation is None or _is_unresolved(annotation): - return "unsupported" - if _is_path_or_str(annotation): - return "scalar" - members = [member for member in _union_members(annotation) if member is not type(None)] - kinds = {_member_input_kind(member, in_union=len(members) > 1) for member in members} - precedence: tuple[ParameterInputKind, ...] = ("unsupported", "collection", "scalar", "in_process_only") - return next((kind for kind in precedence if kind in kinds), "unsupported") - - -def _convert_json_value(annotation: Any, value: Any, *, allow_path: bool) -> Any: - """ - Convert a JSON value to ``annotation``. - - A union takes the value as its first member that matches. ``allow_path`` is True only - for the parameter's own type, so paths are never read from inside a union or collection, - except for an explicit ``Path | str``. - - Returns: - Any: The converted value. - - Raises: - _JsonMismatchError: If the value does not match the annotation. - """ - if annotation is Any: - return value - all_members = _union_members(annotation) - members = [member for member in all_members if member is not type(None)] - if value is None: - if len(members) < len(all_members): - return None - raise _JsonMismatchError - if allow_path and _is_path_or_str(annotation): - if type(value) is str: - return value - raise _JsonMismatchError - if len(members) == 1: - return _convert_json_member(members[0], value, allow_path=allow_path) - for member in members: - try: - return _convert_json_member(member, value, allow_path=False) - except _JsonMismatchError: - continue - raise _JsonMismatchError - - -def _convert_json_member(annotation: Any, value: Any, *, allow_path: bool) -> Any: - """ - Convert a non-null JSON value to one non-union type. - - Choices match the way ``coerce_value`` matches them, so their display strings are accepted. - - Returns: - Any: The converted value. - - Raises: - _JsonMismatchError: If the value does not match the type. - """ - if annotation is Any: - return value - if annotation in (str, bool): - if type(value) is annotation: - return value - elif annotation is int: - if type(value) is int: - return value - elif annotation is float: - if type(value) in (int, float): - try: - return float(value) - except OverflowError: - pass - elif annotation is Path: - if allow_path and type(value) is str: - return Path(value) - elif get_origin(annotation) is Literal or _is_enum_type(annotation): - if type(value) in (str, int, float, bool): - try: - return _coerce_simple_value(param_name="", annotation=annotation, raw_value=value) - except ValueError: - pass - else: - return _convert_json_collection(annotation, value) - raise _JsonMismatchError - - -def _convert_json_collection(annotation: Any, value: Any) -> Any: - """ - Convert a JSON array or object to a declared collection type. - - Returns: - Any: A list, tuple, set, frozenset, or dict holding the converted items. - - Raises: - _JsonMismatchError: If the value does not match the collection type. + bool: True for ``str``/``int``/``float``/``bool``/``Literal``/``Enum`` or a ``list`` of them. """ - origin = get_origin(annotation) or annotation - args = get_args(annotation) - if origin in _MAPPING_ORIGINS: - key_type, value_type = args if len(args) == 2 else (str, Any) - if type(value) is not dict or key_type not in (str, Any): - raise _JsonMismatchError - return {key: _convert_json_value(value_type, item, allow_path=False) for key, item in value.items()} - if origin not in _SEQUENCE_ORIGINS or type(value) is not list: - raise _JsonMismatchError - if origin is tuple and args and not (len(args) == 2 and args[1] is Ellipsis): - if len(args) != len(value): - raise _JsonMismatchError - return tuple(_convert_json_value(arg, item, allow_path=False) for arg, item in zip(args, value, strict=True)) - element_type = args[0] if args else Any - items = [_convert_json_value(element_type, item, allow_path=False) for item in value] - if origin is tuple: - return tuple(items) - if origin in _SET_ORIGINS: - try: - return frozenset(items) if origin is frozenset else set(items) - except TypeError as exc: - raise _JsonMismatchError from exc - return items + if get_origin(annotation) is list: + type_args = get_args(annotation) + annotation = type_args[0] if len(type_args) == 1 else None + return _is_scalar_param_type(annotation) and annotation is not Path and not _is_path_or_str(annotation) def _coerce_simple_value(*, param_name: str, annotation: Any, raw_value: Any) -> Any: diff --git a/pyrit/registry/registry.py b/pyrit/registry/registry.py index 4328f97484..2376a5fb5e 100644 --- a/pyrit/registry/registry.py +++ b/pyrit/registry/registry.py @@ -720,6 +720,37 @@ def create_instance(self, name: str, **kwargs: object) -> T: ) return cls(**resolved) + def create_instance_from_external_input(self, name: str, *, params: Mapping[str, object]) -> T: + """ + Build a configured instance from external input such as a REST or CLI request. + + Unlike ``create_instance``, which serves in-process callers that may pass any + Python object, this accepts only parameters with ``Parameter.is_external_input`` + and registry references given by name. Anything else is rejected before + construction; the constructor still validates the values it receives. + + Args: + name (str): The catalog name to build. + params (Mapping[str, object]): Constructor arguments from the external caller. + + Returns: + T: The constructed instance. + + Raises: + KeyError: If the name is not registered. + ValueError: If an argument is not a valid constructor parameter or not an + external input, a registry reference cannot be resolved, or a value + cannot be coerced. + """ + cls = self.get_class(name) + resolved = resolve_constructor_args( + cls=cls, + raw_args=dict(params), + identifier_type=self._identifier_type(), + external_input=True, + ) + return cls(**resolved) + def __contains__(self, name: str) -> bool: """ Check if a name is registered. @@ -794,6 +825,7 @@ def create_named_instance( type_name: str, params: Mapping[str, object] | None = None, registry_metadata: dict[str, Any] | None = None, + external_input: bool = False, ) -> InstanceT: """ Build and store a configured instance under an explicit name. @@ -804,12 +836,20 @@ def create_named_instance( params (Mapping[str, object] | None): Constructor arguments. registry_metadata (dict[str, Any] | None): Per-entry metadata to store with the instance. + external_input (bool): Whether ``params`` come from an external caller, in which + case the instance is built with ``create_instance_from_external_input``. + Defaults to False. Returns: InstanceT: The constructed and registered instance. """ self.instances.validate_name_available(name) - instance = self.create_instance(type_name, **dict(params) if params is not None else {}) + args = dict(params) if params is not None else {} + instance = ( + self.create_instance_from_external_input(type_name, params=args) + if external_input + else self.create_instance(type_name, **args) + ) self.instances.register(instance, name=name, metadata=registry_metadata) return instance diff --git a/pyrit/registry/registry_metadata.py b/pyrit/registry/registry_metadata.py index 8a6b0c0369..24d21fd5fb 100644 --- a/pyrit/registry/registry_metadata.py +++ b/pyrit/registry/registry_metadata.py @@ -51,11 +51,6 @@ class RegistryMetadata: parameters: tuple[Parameter, ...] = field(kw_only=True, default=()) class_attributes: Mapping[str, Any] = field(kw_only=True, default_factory=dict) - @property - def constructible(self) -> bool: - """Whether registry callers can supply every required parameter as JSON (see ``Parameter.input_kind``).""" - return all(parameter.is_json_configurable for parameter in self.parameters if parameter.required) - @staticmethod def description_from_docstring(cls: type, *, fallback: str = "") -> str: """ diff --git a/pyrit/registry/resolution.py b/pyrit/registry/resolution.py index 42992f68ce..26fe9f8e9d 100644 --- a/pyrit/registry/resolution.py +++ b/pyrit/registry/resolution.py @@ -19,10 +19,10 @@ - **Resolve from a constructor** (``resolve_constructor_args``): derive the contract for a class and turn a flat dict of raw arguments into constructor-ready keyword arguments — coercing simple string values via - ``Parameter.coerce_value``, converting other JSON values under the registry - input contract (``Parameter.input_kind`` / ``Parameter.coerce_json_value``), - and resolving registry-reference parameters by name from the owning domain's - registry. Defaults are left to the constructor. + ``Parameter.coerce_value`` and resolving registry-reference parameters by name + from the owning domain's registry. Defaults are left to the constructor. + Callers building from external input (REST, CLI) select the external path, + which accepts only parameters with ``Parameter.is_external_input``. - **Resolve from a declared list** (``resolve_declared_params``): the sibling for a component that declares an explicit ``list[Parameter]`` (e.g. a scenario's ``supported_parameters()``). It has no references, coerces every supplied @@ -44,6 +44,7 @@ import re import types from collections.abc import Collection, Sequence +from enum import Enum from typing import TYPE_CHECKING, Any, Protocol, TypeAlias, Union, get_args, get_origin, get_type_hints from pydantic import TypeAdapter, ValidationError @@ -58,7 +59,7 @@ from pyrit.models.parameter import display_choices as display_choices if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Mapping from pyrit.models.identifiers.component_identifier import ComponentIdentifier @@ -410,8 +411,7 @@ def _resolve_single_reference( Resolve a single registry-reference value to a stored instance. A string value is looked up by name in the paired registry. An already-built - instance passes through unchanged. Other JSON data (a number or an object) can - never name an instance, so it is rejected instead of reaching the constructor. + instance passes through unchanged. Args: value (Any): The raw value (a registry name, or an instance to pass through). @@ -423,15 +423,9 @@ def _resolve_single_reference( Any: The resolved instance. Raises: - ValueError: If the name is not registered, or the value is JSON data rather than a name - or instance. + ValueError: If the name is not registered. """ if not isinstance(value, str): - if value is not None and _is_json_data(value): - raise ValueError( - f"{owner}.{name}: expected a registry name or instance for this reference, " - f"but got {type(value).__name__}." - ) return value registry = getter() @@ -486,9 +480,8 @@ def _resolve_registry_reference( Any: The resolved instance, or a list of resolved instances. Raises: - ValueError: If a name is not registered, a value is JSON data rather than a - name or instance, or the value's shape (list vs. scalar) does not match the - reference's arity. + ValueError: If a name is not registered, or the value's shape (list vs. + scalar) does not match the reference's arity. """ if get_origin(annotation) is list: if not isinstance(value, list): @@ -531,8 +524,7 @@ def resolve_reference_value( Any: The resolved instance, or the value unchanged when already an instance. Raises: - ValueError: If no registry is wired for ``component_type``, the name is not registered, - or the value is JSON data rather than a name or instance. + ValueError: If no registry is wired for ``component_type``, or the name is not registered. """ getter = _registry_getter_for_component_type(component_type) if getter is None: @@ -545,17 +537,15 @@ def resolve_constructor_args( cls: type, raw_args: dict[str, Any], identifier_type: type[ComponentIdentifier] | None = None, + external_input: bool = False, ) -> dict[str, Any]: """ Resolve a flat argument dict into constructor-ready keyword arguments. Derives the ``Parameter`` contract for ``cls`` and applies it to ``raw_args``. For each raw argument: validate it is a declared parameter; - resolve registry-reference parameters by name; build structured inputs from - their declared variants; coerce simple string values via - ``Parameter.coerce_value``; convert other JSON values via - ``Parameter.coerce_json_value``, which rejects JSON for parameters that take - only Python objects; pass live Python objects through unchanged. + resolve registry-reference parameters by name; coerce simple string values + via ``Parameter.coerce_value``; pass everything else through unchanged. Args: cls (type): The class being built. @@ -563,16 +553,22 @@ def resolve_constructor_args( identifier_type (type[ComponentIdentifier] | None): The domain identifier whose ``Param.*`` markers declare which parameters are registry references. When None, no parameter is treated as a reference. + external_input (bool): Whether ``raw_args`` come from an external caller (REST, CLI). + External callers may set only parameters with ``Parameter.is_external_input`` + and must name registry references. Defaults to False (in-process callers, which + may pass any Python object). Returns: dict[str, Any]: Arguments ready to pass to ``cls(**resolved)``. Raises: ValueError: If an argument is not a declared parameter, a registry - reference cannot be resolved, a simple value cannot be coerced, or a - JSON value is not accepted for the parameter. + reference cannot be resolved, a simple value cannot be coerced, or + external input sets a parameter that is not an external input. """ by_name = {param.name: param for param in derive_parameters(cls=cls, identifier_type=identifier_type)} + if external_input: + reject_non_external_params(params=raw_args, declared=list(by_name.values()), owner=cls.__name__) resolved: dict[str, Any] = {} for name, value in raw_args.items(): @@ -582,6 +578,7 @@ def resolve_constructor_args( f"Unknown parameter '{name}' for '{cls.__name__}'. Valid parameters: {sorted(by_name.keys())}" ) + value_type = _unwrap_optional(param.param_type) if param.reference is not None: getter = _registry_getter_for_component_type(param.reference.component_type) if getter is None: @@ -598,44 +595,61 @@ def resolve_constructor_args( ) elif param.variants is not None: resolved[name] = _resolve_structured_input(parameter=param, value=value) - elif isinstance(value, str) and param.is_string_coercible: + elif (isinstance(value, str) and param.is_string_coercible) or ( + isinstance(value_type, type) and issubclass(value_type, Enum) + ): try: resolved[name] = param.coerce_value(value) except (ValueError, TypeError) as e: raise ValueError(f"Parameter '{name}' of '{cls.__name__}': {e}") from e - elif _is_json_data(value): - resolved[name] = param.coerce_json_value(value, owner=cls.__name__) else: resolved[name] = value return resolved -def _is_json_data(value: Any) -> bool: +def reject_non_external_params(*, params: Mapping[str, Any], declared: Sequence[Parameter], owner: str) -> None: """ - Return whether a value holds only JSON data rather than live Python objects. + Reject external input that is not an explicitly supported external input. - Returns: - bool: True for None, strings, numbers, booleans, and lists or string-keyed dicts of - them (exact built-in types, so subclasses and enums count as live objects), without - cycles. - """ - pending = [value] - seen: set[int] = set() - while pending: - item = pending.pop() - if item is None or type(item) in (str, int, float, bool): - continue - if id(item) in seen or type(item) not in (dict, list): - return False - seen.add(id(item)) - if type(item) is dict: - if not all(type(key) is str for key in item): - return False - pending.extend(item.values()) - else: - pending.extend(item) - return True + Every name must be declared and have ``Parameter.is_external_input``, and registry + references must be given by name. Values are otherwise left to the component. + + Args: + params (Mapping[str, Any]): The parameter values supplied by an external caller. + declared (Sequence[Parameter]): The parameters the component declares. + owner (str): The owning class or scenario name, for error messages. + + Raises: + ValueError: If ``params`` sets an undeclared parameter or one that is not an external + input, or gives a registry reference as anything but a name. + """ + declared_by_name = {parameter.name: parameter for parameter in declared} + for name, value in params.items(): + parameter = declared_by_name.get(name) + if parameter is None: + external_names = sorted( + known for known, candidate in declared_by_name.items() if candidate.is_external_input + ) + raise ValueError(f"Unknown parameter '{name}' for '{owner}'. Valid parameters: {external_names}") + if not parameter.is_external_input: + raise ValueError( + f"Parameter '{name}' of '{owner}' cannot be set through the API; pass it from Python instead." + ) + if parameter.reference is not None: + _require_reference_names(value=value, owner=owner, name=name) + + +def _require_reference_names(*, value: Any, owner: str, name: str) -> None: + """ + Require external input for a registry reference to be a registry name or a list of names. + + Raises: + ValueError: If the value is neither null, a name, nor a list of names. + """ + names = value if isinstance(value, list) else [value] + if value is not None and not all(isinstance(item, str) for item in names): + raise ValueError(f"{owner}.{name}: expected a registry name, but got {type(value).__name__}.") def _resolve_structured_input(*, parameter: Parameter, value: Any) -> Any: diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index 812080b3f9..0d91115ec1 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -213,21 +213,17 @@ async def test_list_converter_types_includes_supported_types(self) -> None: assert "text" in base64_entry.supported_input_types assert "text" in base64_entry.supported_output_types - async def test_types_include_every_registered_converter(self) -> None: - """The projection surfaces every registered converter, including base/helper classes. - - Whether to display a given converter is left to the caller (e.g. the frontend), - so the service hides nothing and reports which ones the API can construct. - """ + async def test_types_omit_converters_that_need_python_objects(self) -> None: + """Converters whose required parameters take Python objects cannot be built from the API.""" service = ConverterService() result = await service.list_converter_types_async() - constructible = {item.converter_type: item.constructible for item in result.items} - assert constructible["Base64Converter"] is True - assert constructible["SearchReplaceConverter"] is True - assert constructible["SelectiveTextConverter"] is False - assert constructible["TextJailbreakConverter"] is False + converter_types = [item.converter_type for item in result.items] + assert "Base64Converter" in converter_types + assert "SearchReplaceConverter" in converter_types + assert "SelectiveTextConverter" not in converter_types + assert "TextJailbreakConverter" not in converter_types async def test_types_serialize_parameter_type(self) -> None: """Type entries render the raw annotation into a human-readable type_name.""" @@ -261,15 +257,24 @@ async def test_types_include_registry_reference_params(self) -> None: target_param = next(param for param in persuasion_entry.parameters if param.name == "converter_target") assert target_param.reference_type == "target" - async def test_types_preserve_all_registry_parameters(self, upload_service: ConverterService) -> None: + async def test_types_expose_only_external_inputs(self, upload_service: ConverterService) -> None: result = await upload_service.list_converter_types_async() metadata_by_name = { metadata.class_name: metadata for metadata in upload_service._registry.get_all_registered_class_metadata() } + expected = { + name + for name, metadata in metadata_by_name.items() + if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) + } - assert {entry.converter_type for entry in result.items} == set(metadata_by_name) + assert {entry.converter_type for entry in result.items} == expected for entry in result.items: - assert entry.parameters == list(metadata_by_name[entry.converter_type].parameters) + assert entry.parameters == [ + parameter + for parameter in metadata_by_name[entry.converter_type].parameters + if parameter.is_external_input + ] @pytest.mark.parametrize( ("converter_type", "parameter_name", "type_name", "required", "is_list"), @@ -415,15 +420,26 @@ async def test_create_converter_raises_for_invalid_type(self) -> None: with pytest.raises(ValueError, match="not found"): await service.create_converter_async(request=request) - async def test_create_converter_rejects_wrong_parameter_type(self) -> None: - """A JSON value that does not match the declared type is a validation error, not a crash.""" + async def test_create_converter_rejects_parameters_that_take_python_objects(self) -> None: service = ConverterService() - request = CreateConverterRequest(name="caesar", type="CaesarConverter", params={"caesar_offset": [1]}) + request = CreateConverterRequest( + name="jailbreak", type="TextJailbreakConverter", params={"jailbreak_template": {"name": "x"}} + ) - with pytest.raises(ValueError, match="caesar_offset"): + with pytest.raises(ValueError, match="'jailbreak_template' of 'TextJailbreakConverter' cannot be set"): await service.create_converter_async(request=request) - assert service.get_converter_object(converter_id="caesar") is None + assert service.get_converter_object(converter_id="jailbreak") is None + + async def test_create_converter_accepts_string_for_string_union_parameter(self) -> None: + service = ConverterService() + request = CreateConverterRequest( + name="replace", type="SearchReplaceConverter", params={"pattern": "a", "replace": "b"} + ) + + result = await service.create_converter_async(request=request) + + assert result.converter_id == "replace" async def test_create_converter_success(self) -> None: """Test successful converter creation.""" @@ -779,9 +795,7 @@ async def test_persist_data_uri_rejects_non_base64_data_uri(self, upload_service assert list(service._upload_path.iterdir()) == [] - async def test_create_converter_cleans_upload_when_construction_fails( - self, upload_service: ConverterService - ) -> None: + async def test_create_converter_cleans_upload_when_creation_fails(self, upload_service: ConverterService) -> None: service = upload_service params = { "existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n"), @@ -789,7 +803,7 @@ async def test_create_converter_cleans_upload_when_construction_fails( } request = CreateConverterRequest(name="invalid-pdf", type="PDFConverter", params=params) - with pytest.raises(ValueError, match="Invalid font_color"): + with pytest.raises(ValueError, match="'font_color' of 'PDFConverter' cannot be set through the API"): await service.create_converter_async(request=request) assert service._registry.instances.get("invalid-pdf") is None diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index cb12bf8f6a..2fe63413af 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -57,6 +57,7 @@ ) from pyrit.models.catalog.scenario import RunScenarioRequest, ScenarioTechniqueSummary from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.registry import ScenarioRegistry from pyrit.scenario import Scenario from pyrit.scenario.core import ( DatasetAttackConfiguration, @@ -64,6 +65,7 @@ get_default_adversarial_target, ) from pyrit.scenario.core.scenario_technique import ScenarioTechnique +from pyrit.scenario.scenarios.airt.scam import Scam from pyrit.score.scorer_evaluation.scorer_metrics import ObjectiveScorerMetrics from unit.mocks import MockPromptTarget, get_mock_target_identifier, make_scenario_result @@ -159,6 +161,37 @@ def _make_request( ) +@pytest.mark.parametrize( + ("scenario_params", "rejected"), + [ + ({"dataset_config": "x"}, "'dataset_config' of 'airt.scam' cannot be set through the API"), + ({"objective_target": {"name": "x"}}, "airt.scam.objective_target: expected a registry name"), + ({"unknown": None}, "Unknown parameter 'unknown' for 'airt.scam'"), + ({"max_concurrency": 2}, None), + ], +) +async def test_start_run_accepts_only_external_scenario_params( + patch_central_database: MagicMock, scenario_params: dict[str, Any], rejected: str | None +) -> None: + registry = MagicMock(spec=ScenarioRegistry) + registry.__contains__.return_value = True + registry.get_class.return_value = Scam + service = ScenarioRunService() + request = _make_request(scenario_name="airt.scam", scenario_params=scenario_params) + + with ( + patch(f"{_REGISTRY_PATCH_BASE}.ScenarioRegistry.get_registry_singleton", return_value=registry), + patch.object(service, "_start_run_locked_async", new_callable=AsyncMock) as start, + ): + if rejected: + with pytest.raises(ValueError, match=rejected): + await service.start_run_async(request=request) + else: + await service.start_run_async(request=request) + + assert start.await_count == (0 if rejected else 1) + + def _make_db_scenario_result( *, result_id: str = "sr-uuid-1", diff --git a/tests/unit/backend/test_scenario_service.py b/tests/unit/backend/test_scenario_service.py index 4b3551baa3..5a32e7b224 100644 --- a/tests/unit/backend/test_scenario_service.py +++ b/tests/unit/backend/test_scenario_service.py @@ -8,6 +8,7 @@ import asyncio import threading from collections import OrderedDict +from dataclasses import replace from typing import TYPE_CHECKING, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -133,6 +134,44 @@ def _make_scenario_metadata( ) +def test_catalog_lists_only_external_scenario_parameters() -> None: + metadata = replace( + _make_scenario_metadata(), + supported_parameters=( + Parameter(name="dataset_config", description="Dataset source configuration.", opaque=True), + Parameter(name="max_concurrency", description="Maximum concurrency.", param_type=int, default=4), + ), + ) + + summary = _metadata_to_registered_scenario(metadata=metadata) + + assert [parameter.name for parameter in summary.supported_parameters] == ["max_concurrency"] + + +@pytest.mark.parametrize( + ("scenario_params", "message"), + [ + ({"dataset_config": "x"}, "'dataset_config' of 'airt.scam' cannot be set through the API"), + ({"objective_target": {"name": "x"}}, "airt.scam.objective_target: expected a registry name"), + ({"unknown": None}, "Unknown parameter 'unknown' for 'airt.scam'"), + ], +) +async def test_configured_estimate_rejects_unsupported_scenario_params_async( + scenario_params: dict[str, object], message: str +) -> None: + registry = MagicMock(spec=ScenarioRegistry) + registry.get_class.return_value = Scam + with patch.object(ScenarioRegistry, "get_registry_singleton", return_value=registry): + service = ScenarioService() + with pytest.raises(ValueError, match=message): + await service.estimate_scenario_run_size_async( + scenario_name="airt.scam", + request=ScenarioRunSizeEstimateRequest(scenario_params=scenario_params), + ) + + registry.create_and_estimate_async.assert_not_called() + + @pytest.mark.parametrize("uses_default", [False, True]) def test_catalog_preserves_adversarial_default_usage(uses_default: bool) -> None: """The public catalog exposes usage from real registry metadata.""" @@ -1090,6 +1129,9 @@ async def estimate_async( service = ScenarioService() service._registry = MagicMock() service._registry.get_registered_class_metadata.return_value = metadata + service._registry.get_class.return_value.supported_parameters.return_value = [ + Parameter(name=name, description="", param_type=int) for name in ("first", "second") + ] service._estimate_configured_run_size_async = AsyncMock(side_effect=estimate_async) first = asyncio.create_task( service.estimate_scenario_run_size_async( @@ -1247,6 +1289,9 @@ async def estimate_async( service = ScenarioService() service._registry = MagicMock() service._registry.get_registered_class_metadata.return_value = metadata + service._registry.get_class.return_value.supported_parameters.return_value = [ + Parameter(name="request_index", description="", param_type=int) + ] service._configured_estimate_semaphore = asyncio.Semaphore(2) service._estimate_configured_run_size_async = AsyncMock(side_effect=estimate_async) @@ -1447,6 +1492,10 @@ async def test_configured_estimate_uses_shared_launch_resolution(self) -> None: introspection_instance._technique_class = _EstimateTechnique introspection_instance._default_dataset_config = DatasetAttackConfiguration(dataset_names=["harmbench"]) scenario_class = MagicMock(return_value=introspection_instance) + scenario_class.supported_parameters.return_value = [ + Parameter(name=name, description="", param_type=int) + for name in ("num_jailbreaks", "num_jailbreak_attempts") + ] objective_target = MagicMock() with ( diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index 5477ac8e1f..4f741b9a7e 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -247,16 +247,6 @@ async def test_types_return_known_target_types(self) -> None: assert "OpenAIChatTarget" in target_types assert "AzureMLChatTarget" in target_types - async def test_types_report_whether_required_parameters_can_be_supplied(self) -> None: - service = TargetService() - - result = await service.list_target_types_async() - - constructible = {item.target_type: item.constructible for item in result.items} - assert constructible["OpenAIChatTarget"] is True - assert constructible["WebsocketTarget"] is False - assert constructible["PlaywrightTarget"] is False - async def test_types_include_declarative_auth_facts(self) -> None: """Type entries surface the per-class auth facts the frontend needs.""" service = TargetService() @@ -282,16 +272,32 @@ async def test_types_include_structured_parameters(self) -> None: assert weights_parameter.is_list is True assert weights_parameter.required is False - async def test_types_preserve_all_registry_parameters(self) -> None: + async def test_types_expose_only_external_inputs(self) -> None: service = TargetService() result = await service.list_target_types_async() metadata_by_name = { metadata.class_name: metadata for metadata in service._registry.get_all_registered_class_metadata() } + expected = { + name + for name, metadata in metadata_by_name.items() + if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) + } - assert {entry.target_type for entry in result.items} == set(metadata_by_name) + assert {entry.target_type for entry in result.items} == expected for entry in result.items: - assert entry.parameters == list(metadata_by_name[entry.target_type].parameters) + assert entry.parameters == [ + parameter for parameter in metadata_by_name[entry.target_type].parameters if parameter.is_external_input + ] + + async def test_types_keep_string_api_key_and_omit_object_only_targets(self) -> None: + service = TargetService() + result = await service.list_target_types_async() + + entries = {entry.target_type: entry for entry in result.items} + assert "api_key" in {parameter.name for parameter in entries["OpenAIChatTarget"].parameters} + assert "custom_configuration" not in {parameter.name for parameter in entries["OpenAIChatTarget"].parameters} + assert not {"PlaywrightTarget", "PlaywrightCopilotTarget", "WebsocketTarget"} & set(entries) async def test_types_cold_and_warm_results_are_equal(self) -> None: service = TargetService() @@ -339,13 +345,6 @@ async def test_types_refresh_after_runtime_class_registration(self) -> None: False, ["text/plain", "text/html"], ), - ( - "PlaywrightCopilotTarget", - "copilot_type", - "CopilotType", - False, - ["consumer", "m365"], - ), ], ) async def test_types_include_enum_parameters( @@ -371,6 +370,15 @@ async def test_types_include_enum_parameters( class TestCreateTarget: """Tests for TargetService.create_target method.""" + async def test_create_target_rejects_parameters_that_take_python_objects(self, sqlite_instance) -> None: + service = TargetService() + request = CreateTargetRequest(name="text", type="TextTarget", params={"custom_configuration": {}}) + + with pytest.raises(ValueError, match="'custom_configuration' of 'TextTarget' cannot be set through the API"): + await service.create_target_async(request=request) + + assert service.get_target_object(target_registry_name="text") is None + async def test_create_target_raises_for_invalid_type(self) -> None: """Test that create_target raises for invalid target type.""" service = TargetService() @@ -426,12 +434,16 @@ async def test_create_target_rejects_reserved_route_name(self, sqlite_instance, ) async def test_create_target_delegates_construction_to_registry(self, sqlite_instance) -> None: - """Every target construction path is owned by the registry.""" + """Every target construction path is owned by the registry, through its external-input path.""" service = TargetService() - with patch.object(service._registry, "create_instance", wraps=service._registry.create_instance) as create: + with patch.object( + service._registry, + "create_instance_from_external_input", + wraps=service._registry.create_instance_from_external_input, + ) as create: await service.create_target_async(request=CreateTargetRequest(type="TextTarget", params={})) - create.assert_called_once() + create.assert_called_once_with("TextTarget", params={}) async def test_create_gandalf_target_coerces_level_string(self, sqlite_instance) -> None: """A Gandalf level from the JSON request is coerced to its enum before construction.""" diff --git a/tests/unit/converter/test_word_doc_converter.py b/tests/unit/converter/test_word_doc_converter.py index e715ae6b27..971a4c1b8e 100644 --- a/tests/unit/converter/test_word_doc_converter.py +++ b/tests/unit/converter/test_word_doc_converter.py @@ -10,7 +10,7 @@ from pyrit.converter import ConverterResult, WordDocConverter from pyrit.models import SeedPrompt -from pyrit.registry.resolution import derive_parameters, resolve_constructor_args +from pyrit.registry.resolution import derive_parameters @pytest.fixture @@ -172,12 +172,6 @@ def test_existing_docx_is_declared_as_path_parameter() -> None: assert parameter.is_path -def test_prompt_template_json_value_is_rejected() -> None: - """The registry must resolve ``prompt_template`` so JSON values are checked before construction.""" - with pytest.raises(ValueError, match="prompt_template"): - resolve_constructor_args(cls=WordDocConverter, raw_args={"prompt_template": "hello"}) - - def test_build_identifier_without_template() -> None: """_build_identifier should return correct params when no template is set.""" converter = WordDocConverter() diff --git a/tests/unit/models/test_parameter.py b/tests/unit/models/test_parameter.py index cf7e8421c8..2c2f0fc8cb 100644 --- a/tests/unit/models/test_parameter.py +++ b/tests/unit/models/test_parameter.py @@ -3,10 +3,10 @@ """Unit tests for the unified Parameter model and its coercion methods.""" -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Sequence from enum import Enum from pathlib import Path -from typing import Any, Literal, Protocol, Union +from typing import Any, Literal, Union import pytest from pydantic import ValidationError @@ -87,7 +87,6 @@ def test_scalar_with_default(self) -> None: "choices": None, "is_list": False, "reference_type": None, - "input_kind": "scalar", "variants": None, } @@ -525,153 +524,62 @@ def __init__(self, *, value=None) -> None: assert param.coerce_value(raw) == expected -class _Greeter(Protocol): - def greet(self) -> str: ... - - -class TestInputKind: - """``input_kind`` states how registry callers supply each parameter.""" +class TestIsExternalInput: + """``is_external_input`` marks the parameters REST, CLI, and GUI callers may supply.""" @pytest.mark.parametrize( - ("param_type", "expected"), + "param_type", [ - (str, "scalar"), - (int | None, "scalar"), - (Path, "scalar"), - (Path | str | None, "scalar"), - (Literal["a", "b"], "scalar"), - (_Speed, "scalar"), - (_Speed | Path, "scalar"), - (_Unsupported | str, "scalar"), - (list[str], "collection"), - (tuple[int, int], "collection"), - (set[str], "collection"), - (Sequence[str] | None, "collection"), - (dict[str, Any], "collection"), - (Mapping[str, list[str]], "collection"), - (list[dict[str, Any]], "collection"), - (str | list[str], "collection"), - (_Unsupported, "in_process_only"), - (Callable[[str], str], "in_process_only"), - (_Greeter | None, "in_process_only"), - (list[Path], "in_process_only"), - (Sequence[Path | str], "in_process_only"), - (dict[int, str], "in_process_only"), - (list[str | bytes], "in_process_only"), - (None, "unsupported"), - (Any, "unsupported"), - ("SeedPrompt | None", "unsupported"), - (list["SeedPrompt"], "unsupported"), + str, + int | None, + Path, + Path | str, + Literal["a", "b"], + _Speed, + list[str], + list[_Speed] | None, + str | list[str], + str | Callable[[], str] | None, + _Speed | str, ], ) - def test_value_parameter_kinds(self, param_type: object, expected: str) -> None: - parameter = Parameter(name="p", description="d", param_type=param_type) - - assert parameter.input_kind == expected - assert parameter.model_dump()["input_kind"] == expected - assert parameter.is_json_configurable is (expected in ("scalar", "collection")) + def test_supported_external_types(self, param_type: object) -> None: + assert Parameter(name="p", description="d", param_type=param_type).is_external_input @pytest.mark.parametrize( - ("param_type", "expected"), + "param_type", [ - (int, "scalar"), - (list[str], "collection"), - (tuple[int, int], "collection"), - (_Unsupported, "in_process_only"), - ("SeedPrompt | None", "unsupported"), + None, + Any, + "SeedPrompt | None", + _Unsupported, + Callable[[], str], + tuple[int, int], + dict[str, str], + Sequence[str], + list[Path], + list[Path | str], + int | tuple[int, int], + _Speed | Path, + str | Path | int, + str | dict[str, str], + str | _Unsupported, + str | list[Path], + list[list[str]], + int | Literal["4", "8"], ], ) - def test_input_kind_survives_wire_round_trip(self, param_type: object, expected: str) -> None: - restored = Parameter.model_validate(Parameter(name="p", description="d", param_type=param_type).model_dump()) - - assert restored.input_kind == expected - assert restored.model_dump()["input_kind"] == expected + def test_other_types_take_python_objects_only(self, param_type: object) -> None: + assert not Parameter(name="p", description="d", param_type=param_type).is_external_input - def test_reference_structured_and_opaque_kinds(self) -> None: + def test_references_and_structured_inputs_are_external(self) -> None: reference = Parameter( name="t", description="d", reference=RegistryReference(component_type=ComponentType.TARGET) ) structured = Parameter(name="s", description="d", param_type=_Unsupported, variants={"one": []}) - opaque = Parameter(name="o", description="d", param_type=str, opaque=True) - - assert (reference.input_kind, structured.input_kind, opaque.input_kind) == ( - "reference", - "structured", - "in_process_only", - ) - assert reference.is_json_configurable and structured.is_json_configurable - assert not opaque.is_json_configurable - - -class TestCoerceJsonValue: - """``coerce_json_value`` applies the registry contract to JSON input.""" - - @pytest.mark.parametrize( - ("param_type", "value", "expected"), - [ - (int, 3, 3), - (float, 2, 2.0), - (str | None, None, None), - (_Speed, "fast", _Speed.FAST), - (Literal[1, 2], "2", 2), - (Path, "/data/input.png", Path("/data/input.png")), - (Path | str, "relative.png", "relative.png"), - (tuple[int, int], [1, 2], (1, 2)), - (tuple[str, ...], ["a", "b"], ("a", "b")), - (set[str], ["a", "a"], {"a"}), - (frozenset[int], [1], frozenset({1})), - (list[_Speed], ["slow"], [_Speed.SLOW]), - (dict[str, list[str]], {"k": ["v"]}, {"k": ["v"]}), - (str | list[str], "x", "x"), - (str | list[str], ["x"], ["x"]), - (_Speed | Path, "slow", _Speed.SLOW), - (_Unsupported | str, "text", "text"), - ], - ) - def test_converts_matching_json(self, param_type: object, value: object, expected: object) -> None: - parameter = Parameter(name="p", description="d", param_type=param_type) - - converted = parameter.coerce_json_value(value, owner="Owner") - - assert converted == expected - assert type(converted) is type(expected) - - @pytest.mark.parametrize( - ("param_type", "value"), - [ - (int, True), - (int, 1.5), - (bool, 1), - (str, 5), - (float, "1.5"), - (int, None), - (tuple[int, int], [1]), - (list[str], "abc"), - (list[Path], ["/etc/hostname"]), - (_Speed | Path, "/etc/hostname"), - (dict[str, int], {"k": "v"}), - (set[str], [["nested"]]), - (list[_Speed], ["bogus"]), - (float, 10**400), - ], - ) - def test_rejects_mismatched_json(self, param_type: object, value: object) -> None: - parameter = Parameter(name="p", description="d", default=REQUIRED_VALUE, param_type=param_type) - - with pytest.raises(ValueError, match="Parameter 'p' of 'Owner'"): - parameter.coerce_json_value(value, owner="Owner") - - @pytest.mark.parametrize( - ("param_type", "message"), - [(_Unsupported, "accepts only a Python object"), (Any, "does not support"), (None, "does not support")], - ) - def test_rejects_json_for_parameters_without_json_input(self, param_type: object, message: str) -> None: - parameter = Parameter(name="p", description="d", param_type=param_type) - - with pytest.raises(ValueError, match=message): - parameter.coerce_json_value({"a": 1}, owner="Owner") - def test_none_is_accepted_when_default_is_none(self) -> None: - parameter = Parameter(name="p", description="d", default=None, param_type=_Unsupported) + assert reference.is_external_input + assert structured.is_external_input - assert parameter.coerce_json_value(None, owner="Owner") is None + def test_opaque_parameter_is_not_external(self) -> None: + assert not Parameter(name="o", description="d", param_type=str, opaque=True).is_external_input diff --git a/tests/unit/registry/test_converter_registry.py b/tests/unit/registry/test_converter_registry.py index 4253c16f7a..da92bce70e 100644 --- a/tests/unit/registry/test_converter_registry.py +++ b/tests/unit/registry/test_converter_registry.py @@ -206,6 +206,24 @@ def test_create_named_instance_stores_registry_metadata(self, registry: Converte assert entry.instance is converter assert entry.metadata == {"owned_artifact_paths": ["managed.dat"]} + def test_create_instance_from_external_input_rejects_object_parameters(self, registry: ConverterRegistry): + with pytest.raises(ValueError, match="'jailbreak_template' of 'TextJailbreakConverter' cannot be set"): + registry.create_instance_from_external_input( + "TextJailbreakConverter", params={"jailbreak_template": {"template": "x"}} + ) + + def test_create_named_instance_selects_external_input_explicitly(self, registry: ConverterRegistry): + with pytest.raises(ValueError, match="cannot be set through the API"): + registry.create_named_instance( + name="math", type_name="MathObfuscationConverter", params={"rng": None}, external_input=True + ) + assert registry.instances.get("math") is None + + converter = registry.create_named_instance( + name="caesar", type_name="CaesarConverter", params={"caesar_offset": "3"}, external_input=True + ) + assert registry.instances.get("caesar") is converter + @pytest.mark.parametrize("name", ["preview", "types"]) def test_create_named_instance_rejects_reserved_name(self, registry: ConverterRegistry, name: str): with pytest.raises(ValueError, match="reserved"): diff --git a/tests/unit/registry/test_registry_metadata.py b/tests/unit/registry/test_registry_metadata.py index b6152b7890..40007babfe 100644 --- a/tests/unit/registry/test_registry_metadata.py +++ b/tests/unit/registry/test_registry_metadata.py @@ -3,8 +3,6 @@ from dataclasses import dataclass, field -from pyrit.common import REQUIRED_VALUE -from pyrit.models import Parameter from pyrit.registry.registry import _matches_filters from pyrit.registry.registry_metadata import RegistryMetadata @@ -225,23 +223,3 @@ def test_matches_filters_combined_include_and_exclude(self): ) is False ) - - -class TestConstructible: - """``constructible`` reports whether every required parameter takes JSON input.""" - - def test_optional_object_parameter_keeps_class_constructible(self) -> None: - parameters = ( - Parameter(name="count", description="", default=REQUIRED_VALUE, param_type=int), - Parameter(name="handle", description="", default=None, param_type=object), - ) - - assert RegistryMetadata(class_name="C", class_module="m", parameters=parameters).constructible - - def test_required_object_parameter_makes_class_not_constructible(self) -> None: - parameters = ( - Parameter(name="count", description="", default=REQUIRED_VALUE, param_type=int), - Parameter(name="handle", description="", default=REQUIRED_VALUE, param_type=object), - ) - - assert not RegistryMetadata(class_name="C", class_module="m", parameters=parameters).constructible diff --git a/tests/unit/registry/test_resolution.py b/tests/unit/registry/test_resolution.py index 6caa1a4112..2b079793bc 100644 --- a/tests/unit/registry/test_resolution.py +++ b/tests/unit/registry/test_resolution.py @@ -5,20 +5,16 @@ Tests for the shared registry constructor-argument resolution primitive. """ -import contextlib -import json -from collections.abc import Collection -from dataclasses import dataclass +from collections.abc import Callable from enum import Enum -from pathlib import Path -from typing import TYPE_CHECKING, Any, Literal, Protocol +from typing import Any, Literal import pytest from pyrit.common import REQUIRED_VALUE, forward_init_parameters from pyrit.common.apply_defaults import _RequiredValueSentinel from pyrit.models import Message, MessagePiece -from pyrit.models.identifiers import ConverterIdentifier, ScorerIdentifier, TargetIdentifier +from pyrit.models.identifiers import ConverterIdentifier, TargetIdentifier from pyrit.models.parameter import ComponentType from pyrit.prompt_target import PromptTarget from pyrit.registry.components import ConverterRegistry, ScorerRegistry, TargetRegistry @@ -26,12 +22,10 @@ _registry_getter_for_component_type, derive_parameters, display_choices, + reject_non_external_params, resolve_constructor_args, ) -if TYPE_CHECKING: - from pyrit.prompt_target import PromptTarget as _TypeCheckingOnlyTarget - class MockPromptTarget(PromptTarget): """Minimal PromptTarget for registry-resolution tests.""" @@ -146,79 +140,6 @@ def __init__(self, *, targets: list[PromptTarget]) -> None: self.targets = targets -@dataclass -class _Settings: - level: int = 0 - - -class _Provider(Protocol): - def provide(self) -> str: ... - - -class _Sized(Protocol): - def __len__(self) -> int: ... - - -class _Unresolved: - """Helper whose annotation names a type-checking-only import, as many components do.""" - - def __init__(self, *, target: "_TypeCheckingOnlyTarget | None" = None) -> None: - self.target = target - - -class _Handle: - """A live object type that no JSON value can represent.""" - - -class _Bag(list[int]): - """A live list subclass that callers pass as an existing object.""" - - -class _JsonShaped: - """Helper whose constructor takes container and object parameters that JSON callers supply.""" - - def __init__( - self, - *, - color: tuple[int, int, int] = (0, 0, 0), - weights: list[int] | None = None, - extra: dict[str, int] | None = None, - speed: _Speed | None = None, - speeds: list[_Speed] | None = None, - modes: list[Literal["a", "b"] | None] | None = None, - note: str | None = None, - words: Collection[str] | None = None, - settings: _Settings | None = None, - location: _Speed | Path | None = None, - provider: _Provider | None = None, - sized: _Sized | None = None, - options: dict[str, Any] | None = None, - handle: _Handle | None = None, - groups: dict[str, Collection[str]] | None = None, - choices: Collection[str] | _Handle | None = None, - bag: _Bag | None = None, - anything: Collection | None = None, - ) -> None: - self.color = color - self.weights = weights - self.extra = extra - self.speed = speed - self.speeds = speeds - self.modes = modes - self.note = note - self.words = words - self.settings = settings - self.location = location - self.provider = provider - self.sized = sized - self.options = options - self.handle = handle - self.groups = groups - self.choices = choices - self.bag = bag - self.anything = anything - - def _resolve(cls: type, raw_args: dict[str, object], *, identifier_type: type | None = None) -> dict[str, object]: """Resolve ``raw_args`` against the derived parameter contract for ``cls``.""" return resolve_constructor_args(cls=cls, raw_args=raw_args, identifier_type=identifier_type) @@ -340,243 +261,148 @@ def test_unknown_registry_reference_empty_registry_hint(self, empty_target_regis with pytest.raises(ValueError, match="is empty"): _resolve(_NeedsTarget, {"converter_target": "missing"}, identifier_type=ConverterIdentifier) - @pytest.mark.parametrize( - ("cls", "raw_args"), - [ - (_SimpleOnly, {"count": {"value": 1}}), - (_SimpleOnly, {"count": [5]}), - (_SimpleOnly, {"count": True}), - (_SimpleOnly, {"count": 1.5}), - (_SimpleOnly, {"count": None}), - (_SimpleOnly, {"ratio": {"value": 1}}), - (_SimpleOnly, {"flag": 1}), - (_SimpleOnly, {"ratio": 10**400}), - (_JsonShaped, {"note": 5}), - (_JsonShaped, {"weights": [1, "2"]}), - (_JsonShaped, {"extra": "not-an-object"}), - (_JsonShaped, {"color": [1, 2]}), - (_JsonShaped, {"color": None}), - (_JsonShaped, {"words": [1, "a"]}), - (_JsonShaped, {"words": "the"}), - (_JsonShaped, {"groups": {"group": [1]}}), - (_JsonShaped, {"choices": {}}), - (_JsonShaped, {"anything": 7}), - (_JsonShaped, {"modes": ["c"]}), - ], - ) - def test_rejects_json_value_of_wrong_type(self, cls: type, raw_args: dict[str, object]) -> None: - with pytest.raises(ValueError, match="expects"): - _resolve(cls, raw_args) - - @pytest.mark.parametrize( - ("cls", "raw_args"), - [ - (_SimpleOnly, {"ratio": 1}), - (_SimpleOnly, {"count": 3, "flag": False}), - (_JsonShaped, {"weights": [1, 2], "extra": {"a": 1}}), - (_JsonShaped, {"note": None, "settings": None}), - (_JsonShaped, {"words": ["the", "a"]}), - (_JsonShaped, {"groups": {"group": ["one", "two"]}}), - (_JsonShaped, {"choices": ["the", "a"]}), - (_JsonShaped, {"anything": [1, "a"]}), - (_JsonShaped, {"modes": []}), - (_JsonShaped, {"modes": ["a", None]}), - ], - ) - def test_accepts_matching_json_value_unchanged(self, cls: type, raw_args: dict[str, object]) -> None: - assert _resolve(cls, raw_args) == raw_args - - def test_json_array_becomes_tuple_for_tuple_parameter(self) -> None: - assert _resolve(_JsonShaped, {"color": [10, 20, 30]})["color"] == (10, 20, 30) - - def test_rejects_json_for_in_process_only_parameter(self) -> None: - with pytest.raises(ValueError, match="settings.*accepts only a Python object"): - _resolve(_JsonShaped, {"settings": {"level": 1}}) - def test_union_takes_json_only_through_its_json_members(self) -> None: - assert _resolve(_JsonShaped, {"location": "fast"}) == {"location": _Speed.FAST} - with pytest.raises(ValueError, match="location"): - _resolve(_JsonShaped, {"location": "/tmp/x"}) - path = Path("/tmp/x") - assert _resolve(_JsonShaped, {"location": path})["location"] is path +class _Handle: + """A Python object no external caller can supply.""" - def test_live_objects_pass_through_unchecked(self) -> None: - color = object() - weights = [object()] - options = {"timeout": (5.0, 10.0)} - extra: dict[str, object] = {} - extra["self"] = extra - bag = _Bag([1, 2]) - live = {"color": color, "weights": weights, "options": options, "extra": extra, "bag": bag} - resolved = _resolve(_JsonShaped, live) +class _Mixed: + """Helper whose constructor mixes external inputs with Python-object parameters.""" - assert all(resolved[name] is value for name, value in live.items()) + def __init__( + self, + *, + count: int = 1, + words: list[str] | None = None, + key: str | Callable[[], str] | None = None, + handle: _Handle | None = None, + options: dict[str, Any] | None = None, + ) -> None: + self.count = count + self.words = words + self.key = key + self.handle = handle + self.options = options - def test_deeply_nested_json_value_is_rejected(self) -> None: - nested: object = 1 - for _ in range(500): - nested = [nested] - with pytest.raises(ValueError, match="expects"): - _resolve(_SimpleOnly, {"count": nested}) +@pytest.mark.usefixtures("patch_central_database") +class TestExternalInput: + """The external path accepts only supported inputs; the in-process path is unchanged.""" - @pytest.mark.parametrize("raw_args", [{"provider": "anything"}, {"sized": 7}, {"sized": [1, 2]}]) - def test_protocol_parameter_takes_only_live_objects(self, raw_args: dict[str, object]) -> None: - with pytest.raises(ValueError, match="accepts only a Python object"): - _resolve(_JsonShaped, raw_args) + def test_external_accepts_supported_inputs(self) -> None: + raw_args: dict[str, object] = {"count": "3", "words": ["a", "b"], "key": "secret"} - def test_protocol_parameter_accepts_live_object_and_none(self) -> None: - sized = _Bag([1, 2]) + resolved = resolve_constructor_args(cls=_Mixed, raw_args=raw_args, external_input=True) - assert _resolve(_JsonShaped, {"provider": None, "sized": sized}) == {"provider": None, "sized": sized} + assert resolved == {"count": 3, "words": ["a", "b"], "key": "secret"} - def test_unresolved_annotation_rejects_json(self) -> None: - target = MockPromptTarget() + @pytest.mark.parametrize("raw_args", [{"handle": {}}, {"handle": None}, {"options": {"a": 1}}]) + def test_external_rejects_parameters_that_take_python_objects(self, raw_args: dict[str, object]) -> None: + with pytest.raises(ValueError, match="cannot be set through the API"): + resolve_constructor_args(cls=_Mixed, raw_args=raw_args, external_input=True) - with pytest.raises(ValueError, match="does not support"): - _resolve(_Unresolved, {"target": {"a": 1}}) - assert _resolve(_Unresolved, {"target": None}) == {"target": None} - assert _resolve(_Unresolved, {"target": target})["target"] is target + def test_external_still_rejects_unknown_parameters(self) -> None: + with pytest.raises(ValueError, match="Unknown parameter 'nope'"): + resolve_constructor_args(cls=_Mixed, raw_args={"nope": 1}, external_input=True) - def test_json_value_for_object_parameter_is_rejected(self) -> None: + def test_in_process_callers_keep_passing_python_objects(self) -> None: handle = _Handle() + options = {"dtype": object()} + + resolved = _resolve(_Mixed, {"handle": handle, "options": options, "key": print}) + + assert resolved["handle"] is handle + assert resolved["options"] is options + assert resolved["key"] is print + + def test_external_values_of_supported_parameters_reach_the_constructor(self) -> None: + assert resolve_constructor_args(cls=_Mixed, raw_args={"words": "a, b"}, external_input=True) == { + "words": "a, b" + } + + @pytest.mark.parametrize("value", [{"name": "my_target"}, 5, True]) + def test_external_reference_must_be_a_name(self, target_registry: TargetRegistry, value: object) -> None: + with pytest.raises(ValueError, match="expected a registry name"): + resolve_constructor_args( + cls=_NeedsTarget, + raw_args={"converter_target": value}, + identifier_type=ConverterIdentifier, + external_input=True, + ) + + def test_external_list_reference_must_be_names(self, target_registry: TargetRegistry) -> None: + with pytest.raises(ValueError, match="expected a registry name"): + resolve_constructor_args( + cls=_NeedsTargets, + raw_args={"targets": ["my_target", {"name": "x"}]}, + identifier_type=TargetIdentifier, + external_input=True, + ) + + resolved = resolve_constructor_args( + cls=_NeedsTargets, + raw_args={"targets": ["my_target"]}, + identifier_type=TargetIdentifier, + external_input=True, + ) + assert resolved["targets"] == [target_registry.instances.get("my_target")] + + def test_external_reference_resolves_name(self, target_registry: TargetRegistry) -> None: + resolved = resolve_constructor_args( + cls=_NeedsTarget, + raw_args={"converter_target": "my_target"}, + identifier_type=ConverterIdentifier, + external_input=True, + ) - with pytest.raises(ValueError, match="handle"): - _resolve(_JsonShaped, {"handle": {}}) - assert _resolve(_JsonShaped, {"handle": None}) == {"handle": None} - assert _resolve(_JsonShaped, {"handle": handle})["handle"] is handle + assert resolved["converter_target"] is target_registry.instances.get("my_target") - @pytest.mark.parametrize( - ("registry_type", "identifier_type", "type_name", "raw_args", "message"), - [ - (TargetRegistry, TargetIdentifier, "TextTarget", {"custom_configuration": {}}, "accepts only"), - (TargetRegistry, TargetIdentifier, "A2ATarget", {"auth_token": {"token": "x"}}, "expects"), - ( - TargetRegistry, - TargetIdentifier, - "OpenAIResponseTarget", - {"tool_providers": [{"name": "x"}]}, - "accepts only", - ), - (ConverterRegistry, ConverterIdentifier, "TokenBijectionConverter", {"tokenizer": "name"}, "accepts only"), - ( - ConverterRegistry, - ConverterIdentifier, - "TextJailbreakConverter", - {"jailbreak_template": {}}, - "accepts only", - ), - ], - ) - def test_registered_component_rejects_json_for_object_parameter( - self, - registry_type: type, - identifier_type: type, - type_name: str, - raw_args: dict[str, object], - message: str, - ) -> None: - cls = registry_type.get_registry_singleton().get_class(type_name) + def test_reject_non_external_params_checks_names_and_references(self) -> None: + declared = derive_parameters(cls=_Mixed) + reference = derive_parameters(cls=_NeedsTarget, identifier_type=ConverterIdentifier) - with pytest.raises(ValueError, match=message): - _resolve(cls, raw_args, identifier_type=identifier_type) + reject_non_external_params(params={"count": 1, "key": "k"}, declared=declared, owner="demo") + with pytest.raises(ValueError, match="'handle' of 'demo' cannot be set through the API"): + reject_non_external_params(params={"handle": None}, declared=declared, owner="demo") + with pytest.raises(ValueError, match="Unknown parameter 'unknown' for 'demo'"): + reject_non_external_params(params={"unknown": None}, declared=declared, owner="demo") + with pytest.raises(ValueError, match="demo.converter_target: expected a registry name"): + reject_non_external_params(params={"converter_target": {"name": "x"}}, declared=reference, owner="demo") @pytest.mark.parametrize( ("registry_type", "identifier_type", "type_name", "raw_args"), [ - (ConverterRegistry, ConverterIdentifier, "FlipConverter", {"converter_target": {"name": "x"}}), - (TargetRegistry, TargetIdentifier, "RoundRobinTarget", {"targets": [{"name": "x"}]}), + (TargetRegistry, TargetIdentifier, "TextTarget", {"custom_configuration": {}}), + (TargetRegistry, TargetIdentifier, "OpenAIChatTarget", {"httpx_client_kwargs": {"timeout": 5}}), + (ConverterRegistry, ConverterIdentifier, "TextJailbreakConverter", {"jailbreak_template": "x"}), + (ConverterRegistry, ConverterIdentifier, "PDFConverter", {"font_color": [1, 2, 3]}), ], ) - def test_registry_reference_rejects_json_object( + def test_registered_components_reject_object_parameters_from_external_input( self, registry_type: type, identifier_type: type, type_name: str, raw_args: dict[str, object] ) -> None: cls = registry_type.get_registry_singleton().get_class(type_name) - with pytest.raises(ValueError, match="registry name or instance"): - _resolve(cls, raw_args, identifier_type=identifier_type) - - def test_registered_target_accepts_string_token(self) -> None: - cls = TargetRegistry.get_registry_singleton().get_class("A2ATarget") - - assert _resolve(cls, {"auth_token": "token"}, identifier_type=TargetIdentifier) == {"auth_token": "token"} + with pytest.raises(ValueError, match="cannot be set through the API"): + resolve_constructor_args(cls=cls, raw_args=raw_args, identifier_type=identifier_type, external_input=True) @pytest.mark.parametrize( - ("type_name", "raw_args", "expected"), + ("registry_type", "identifier_type", "type_name", "raw_args"), [ - ("ImageCompressionConverter", {"background_color": [10, 20, 30]}, {"background_color": (10, 20, 30)}), - ("SATAMaskingConverter", {"stopwords": ["the", "a"]}, {"stopwords": ["the", "a"]}), + (TargetRegistry, TargetIdentifier, "OpenAIChatTarget", {"api_key": "key", "temperature": 0.5}), + (ConverterRegistry, ConverterIdentifier, "SearchReplaceConverter", {"pattern": "a", "replace": "b"}), ], ) - def test_registered_converter_json_values( - self, type_name: str, raw_args: dict[str, object], expected: dict[str, object] + def test_registered_components_accept_string_inputs_from_external_input( + self, registry_type: type, identifier_type: type, type_name: str, raw_args: dict[str, object] ) -> None: - cls = ConverterRegistry.get_registry_singleton().get_class(type_name) - - assert _resolve(cls, raw_args, identifier_type=ConverterIdentifier) == expected - - @pytest.mark.parametrize( - ("registry_type", "identifier_type"), - [(ConverterRegistry, ConverterIdentifier), (TargetRegistry, TargetIdentifier)], - ) - def test_registered_parameters_only_raise_value_error(self, registry_type: type, identifier_type: type) -> None: - registry = registry_type.get_registry_singleton() - for type_name in registry.get_class_names(): - cls = registry.get_class(type_name) - for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): - for value in (None, 1, 1.5, True, "text", [], [1, "a"], {"a": [1]}): - with contextlib.suppress(ValueError): - _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) + cls = registry_type.get_registry_singleton().get_class(type_name) - @pytest.mark.parametrize( - ("registry_type", "identifier_type"), - [(ConverterRegistry, ConverterIdentifier), (TargetRegistry, TargetIdentifier)], - ) - def test_registered_json_defaults_are_accepted(self, registry_type: type, identifier_type: type) -> None: - registry = registry_type.get_registry_singleton() - for type_name in registry.get_class_names(): - cls = registry.get_class(type_name) - for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): - if parameter.input_kind not in ("scalar", "collection") or parameter.default is None: - continue - try: - value = json.loads(json.dumps(parameter.default)) - except (TypeError, ValueError): - continue - _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) + resolved = resolve_constructor_args( + cls=cls, raw_args=raw_args, identifier_type=identifier_type, external_input=True + ) - @pytest.mark.parametrize( - ("registry_type", "identifier_type"), - [ - (ConverterRegistry, ConverterIdentifier), - (TargetRegistry, TargetIdentifier), - (ScorerRegistry, ScorerIdentifier), - ], - ) - def test_registered_choices_are_accepted(self, registry_type: type, identifier_type: type) -> None: - registry = registry_type.get_registry_singleton() - for type_name in registry.get_class_names(): - cls = registry.get_class(type_name) - for parameter in derive_parameters(cls=cls, identifier_type=identifier_type): - if parameter.reference is not None or not parameter.choices: - continue - value = [parameter.choices[0]] if parameter.is_list else parameter.choices[0] - _resolve(cls, {parameter.name: value}, identifier_type=identifier_type) - - def test_enum_list_is_coerced_from_json_choices(self) -> None: - live = [_Speed.SLOW] - - assert _resolve(_JsonShaped, {"speeds": ["fast"]}) == {"speeds": [_Speed.FAST]} - assert _resolve(_JsonShaped, {"speeds": live})["speeds"] is live - with pytest.raises(ValueError, match="speeds"): - _resolve(_JsonShaped, {"speeds": ["bogus"]}) - - def test_collection_parameter_rejects_wrong_member_type(self) -> None: - cls = ConverterRegistry.get_registry_singleton().get_class("SATAMaskingConverter") - - with pytest.raises(ValueError, match="stopwords"): - _resolve(cls, {"stopwords": [1, "a"]}, identifier_type=ConverterIdentifier) + assert resolved == raw_args class TestDeriveParameters: diff --git a/tests/unit/score/test_garak_exploitation_scorer.py b/tests/unit/score/test_garak_exploitation_scorer.py index 5b9e2da013..48738c8123 100644 --- a/tests/unit/score/test_garak_exploitation_scorer.py +++ b/tests/unit/score/test_garak_exploitation_scorer.py @@ -78,16 +78,11 @@ async def test_default_corpus_is_ready_to_score_async( assert (await scorer.score_text_async(text))[0].get_value() is True assert (await scorer.score_text_async("I cannot help with that request."))[0].get_value() is False - @pytest.mark.parametrize("payloads", [[], [""], b"payload()", ("payload()",)]) + @pytest.mark.parametrize("payloads", [[], [""], "payload()", b"payload()", ("payload()",), [1], ["valid", None]]) def test_invalid_payloads_raise(self, payloads: object) -> None: with pytest.raises(ValueError, match="nonempty list of nonempty strings"): ScorerRegistry().create_instance("GarakExploitationScorer", detector="jinja", payloads=payloads) - @pytest.mark.parametrize("payloads", ["payload()", [1], ["valid", None]]) - def test_mistyped_json_payloads_raise(self, payloads: object) -> None: - with pytest.raises(ValueError, match=r"expects list\[str\]"): - ScorerRegistry().create_instance("GarakExploitationScorer", detector="jinja", payloads=payloads) - async def test_registry_list_contract_async(self) -> None: registry = ScorerRegistry() metadata = registry.get_class_metadata(GarakExploitationScorer) From 3553736b46840dbcaf8a2b25d75d98f0b9178979 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 17:07:50 +0000 Subject: [PATCH 4/8] Accept scalar collections and scalar alternatives, and check scenario params in the worker - Treat flat list/Collection/Sequence of non-path scalars, and unions with such an alternative and no path alternative, as external inputs, so font_size=24, n_seconds=8, and word lists work through the API. - Check scenario_params in the preparation worker after resolving the scenario class, so a cold registry is never discovered on the API event loop. - Build REST-created scorers through the external path and list only the scorer parameters external callers may set. - Add independent catalog and creation tests, real constructor-failure cleanup coverage, and file-collection rejection coverage. --- .../backend/services/scenario_run_service.py | 13 +-- pyrit/backend/services/scorer_service.py | 11 +- pyrit/models/parameter.py | 36 ++++-- tests/unit/backend/test_converter_service.py | 105 +++++++++++++++--- .../unit/backend/test_scenario_run_service.py | 36 ++++-- tests/unit/backend/test_scorer_service.py | 39 ++++++- tests/unit/backend/test_target_service.py | 49 +++++--- tests/unit/models/test_parameter.py | 21 +++- 8 files changed, 241 insertions(+), 69 deletions(-) diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index a062e95832..2f6928efb0 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -244,13 +244,6 @@ async def start_run_async(self, *, request: RunScenarioRequest) -> ScenarioRunSu Returns: ScenarioRunSummary: Current scheduled run state. """ - registry = ScenarioRegistry.get_registry_singleton() - if request.scenario_params and request.scenario_name in registry: - reject_non_external_params( - params=request.scenario_params, - declared=registry.get_class(request.scenario_name).supported_parameters(), - owner=request.scenario_name, - ) async with self._reserve_resume_request_async(request.scenario_result_id), self._launch_lock: await self._validate_resume_admission_async(scenario_result_id=request.scenario_result_id) return await self._start_run_locked_async(request=request) @@ -633,6 +626,12 @@ async def _prepare_run_async(self, *, request: RunScenarioRequest) -> _PreparedR ValueError: If scenario, target, initializer, or technique cannot be found. """ scenario_class = self._configuration_resolver.resolve_scenario_class(scenario_name=request.scenario_name) + if request.scenario_params: + reject_non_external_params( + params=request.scenario_params, + declared=scenario_class.supported_parameters(), + owner=request.scenario_name, + ) await self._run_initializers_async(request=request) objective_target = self._configuration_resolver.resolve_target(target_name=request.target_name) adversarial_target = self._configuration_resolver.resolve_adversarial_target( diff --git a/pyrit/backend/services/scorer_service.py b/pyrit/backend/services/scorer_service.py index ba8b711e44..49328a54f5 100644 --- a/pyrit/backend/services/scorer_service.py +++ b/pyrit/backend/services/scorer_service.py @@ -36,21 +36,25 @@ def _build_instance(self, *, name: str, scorer: Any) -> ScorerInstance: async def list_scorer_types_async(self) -> ScorerTypeResponse: """ - List registered scorer class metadata without constructing scorers. + List the scorer types external callers can build, without constructing scorers. + + Each entry lists only the parameters external callers may supply; types that + need a Python object for a required parameter are left out. Returns: - ScorerTypeResponse: All registered scorer type metadata. + ScorerTypeResponse: Scorer type metadata for external callers. """ def list_types() -> ScorerTypeResponse: items = [ ScorerTypeEntry( scorer_type=metadata.class_name, - parameters=list(metadata.parameters), + parameters=[parameter for parameter in metadata.parameters if parameter.is_external_input], is_llm_based=metadata.is_llm_based, description=metadata.class_description or None, ) for metadata in self._registry.get_all_registered_class_metadata() + if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) ] return ScorerTypeResponse(items=items) @@ -110,6 +114,7 @@ def create() -> ScorerInstance: name=request.name, type_name=request.type, params=request.params, + external_input=True, ) return self._build_instance(name=request.name, scorer=scorer) diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index 782ea2328d..0bf841d7ca 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -8,7 +8,7 @@ import copy import types from abc import ABC, abstractmethod -from collections.abc import Callable +from collections.abc import Collection, Sequence from dataclasses import dataclass from enum import Enum from pathlib import Path @@ -283,11 +283,12 @@ def is_external_input(self) -> bool: """ Whether REST, CLI, and GUI callers may supply this parameter. - True for registry references, declared structured inputs, scalars, lists of - non-path scalars, and unions of ``str`` with those or with callables (external - callers send the string, as for ``api_key: str | Callable[...]``; callables are - for in-process callers). Other parameters take Python objects from in-process - callers only. + True for registry references, declared structured inputs, scalars, flat + ``list`` / ``Collection`` / ``Sequence`` of non-path scalars, and unions with one + of those as an alternative and no path alternative. External callers supply that + alternative, as for ``api_key: str | Callable[...]`` or + ``font_size: int | tuple[int, int]``; the other alternatives are for in-process + callers. Other parameters take Python objects from in-process callers only. Returns: bool: True when external callers may supply this parameter. @@ -301,8 +302,8 @@ def is_external_input(self) -> bool: return True if get_origin(param_type) in (Union, types.UnionType): members = [member for member in get_args(param_type) if member is not type(None)] - return str in members and all( - _is_non_path_json_type(member) or get_origin(member) is Callable for member in members + return not any(_mentions_path(member) for member in members) and any( + _is_non_path_json_type(member) for member in members ) return _is_non_path_json_type(param_type) @@ -442,17 +443,30 @@ def _is_scalar_param_type(annotation: Any) -> bool: def _is_non_path_json_type(annotation: Any) -> bool: """ - Return whether the annotation is a non-path scalar or a ``list`` of non-path scalars. + Return whether the annotation is a non-path scalar or a flat collection of one. + + A flat collection is a ``list``, ``Collection``, or ``Sequence`` of a single non-path + scalar; external callers send it as a JSON array, which reaches the constructor as a list. Returns: - bool: True for ``str``/``int``/``float``/``bool``/``Literal``/``Enum`` or a ``list`` of them. + bool: True for ``str``/``int``/``float``/``bool``/``Literal``/``Enum`` or a flat collection of them. """ - if get_origin(annotation) is list: + if get_origin(annotation) in (list, Collection, Sequence): type_args = get_args(annotation) annotation = type_args[0] if len(type_args) == 1 else None return _is_scalar_param_type(annotation) and annotation is not Path and not _is_path_or_str(annotation) +def _mentions_path(annotation: Any) -> bool: + """ + Return whether the annotation is ``Path`` or has ``Path`` among its type arguments. + + Returns: + bool: True when a value of this type may be a local file path. + """ + return annotation is Path or any(_mentions_path(argument) for argument in get_args(annotation)) + + def _coerce_simple_value(*, param_name: str, annotation: Any, raw_value: Any) -> Any: """ Coerce ``raw_value`` to a scalar ``annotation`` — the shared coercion core. diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index c6d181fb09..a0ffa947c6 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -9,11 +9,13 @@ import base64 import codecs from collections.abc import AsyncGenerator +from io import BytesIO from pathlib import Path from unittest.mock import AsyncMock, MagicMock, call, patch import pytest from fastapi import HTTPException +from PIL import Image from pydantic import ValidationError from pyrit import converter @@ -263,24 +265,33 @@ async def test_types_include_registry_reference_params(self) -> None: target_param = next(param for param in persuasion_entry.parameters if param.name == "converter_target") assert target_param.reference_type == "target" - async def test_types_expose_only_external_inputs(self, upload_service: ConverterService) -> None: + async def test_types_expose_component_external_inputs(self, upload_service: ConverterService) -> None: result = await upload_service.list_converter_types_async() - metadata_by_name = { - metadata.class_name: metadata for metadata in upload_service._registry.get_all_registered_class_metadata() - } - expected = { - name - for name, metadata in metadata_by_name.items() - if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) + parameters = { + entry.converter_type: {parameter.name for parameter in entry.parameters} for entry in result.items } - assert {entry.converter_type for entry in result.items} == expected - for entry in result.items: - assert entry.parameters == [ - parameter - for parameter in metadata_by_name[entry.converter_type].parameters - if parameter.is_external_input - ] + assert { + "GridCompositeConverter", + "SelectiveTextConverter", + "TextJailbreakConverter", + "TokenBijectionConverter", + }.isdisjoint(parameters) + assert "font_size" in parameters["AddImageTextConverter"] + assert {"stopwords", "candidate_words"} <= parameters["SATAMaskingConverter"] + assert {"existing_docx", "placeholder"} <= parameters["WordDocConverter"] + assert "existing_pdf" in parameters["PDFConverter"] + assert "font_color" not in parameters["PDFConverter"] + + async def test_registry_metadata_keeps_parameters_the_api_cannot_set( + self, upload_service: ConverterService + ) -> None: + metadata = upload_service._registry.get_registered_class_metadata("PDFConverter") + assert metadata is not None + registry_parameters = {parameter.name: parameter for parameter in metadata.parameters} + + assert registry_parameters["font_color"].type_name == "tuple[int, int, int]" + assert not registry_parameters["font_color"].is_external_input @pytest.mark.parametrize( ("converter_type", "parameter_name", "type_name", "required", "is_list"), @@ -815,6 +826,70 @@ async def test_create_converter_cleans_upload_when_creation_fails(self, upload_s assert service._registry.instances.get("invalid-pdf") is None assert list(service._upload_path.iterdir()) == [] + async def test_create_converter_cleans_upload_when_constructor_rejects_value( + self, upload_service: ConverterService + ) -> None: + docx_mime_type = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" + request = CreateConverterRequest( + name="invalid-word-doc", + type="WordDocConverter", + params={ + "existing_docx": _make_data_uri(mime_type=docx_mime_type, content=b"PK\x03\x04"), + "placeholder": "", + }, + ) + + with pytest.raises(ValueError, match="Placeholder must be a non-empty string"): + await upload_service.create_converter_async(request=request) + + assert upload_service._registry.instances.get("invalid-word-doc") is None + assert list(upload_service._upload_path.iterdir()) == [] + + async def test_create_converter_accepts_scalar_alternative_of_union(self, upload_service: ConverterService) -> None: + image = BytesIO() + Image.new("RGB", (8, 8)).save(image, format="PNG") + request = CreateConverterRequest( + name="sized-text", + type="AddImageTextConverter", + params={"img_to_add": _make_data_uri(mime_type="image/png", content=image.getvalue()), "font_size": 24}, + ) + + response = await upload_service.create_converter_async(request=request) + + converter = upload_service.get_converter_object(converter_id=response.converter_id) + assert converter._font_size == 24 + + async def test_create_converter_accepts_word_lists(self, upload_service: ConverterService) -> None: + request = CreateConverterRequest( + name="masked", + type="SATAMaskingConverter", + params={"stopwords": ["the", "a"], "candidate_words": ["bomb"]}, + ) + + response = await upload_service.create_converter_async(request=request) + + strategy_params = upload_service.get_converter_object( + converter_id=response.converter_id + )._selection_strategy.get_identifier_params() + assert strategy_params["stopwords"] == ["a", "the"] + assert strategy_params["candidate_words"] == ["bomb"] + + @pytest.mark.parametrize("innocuous_image", ["/etc/hosts", "https://example.com/cat.png"]) + async def test_create_converter_rejects_file_collections( + self, upload_service: ConverterService, innocuous_image: str + ) -> None: + request = CreateConverterRequest( + name="grid", + type="GridCompositeConverter", + params={"innocuous_images": [innocuous_image]}, + ) + + with pytest.raises(ValueError, match="'innocuous_images' of 'GridCompositeConverter' cannot be set"): + await upload_service.create_converter_async(request=request) + + assert upload_service._registry.instances.get("grid") is None + assert list(upload_service._upload_path.iterdir()) == [] + async def test_create_converter_registers_nothing_when_response_mapping_fails( self, upload_service: ConverterService ) -> None: diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index 05f2434b61..1134f4cea1 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -65,6 +65,7 @@ get_default_adversarial_target, ) from pyrit.scenario.core.scenario_technique import ScenarioTechnique +from pyrit.scenario.scenarios.airt.jailbreak import Jailbreak from pyrit.scenario.scenarios.airt.scam import Scam from pyrit.score.scorer_evaluation.scorer_metrics import ObjectiveScorerMetrics from unit.mocks import MockPromptTarget, get_mock_target_identifier, make_scenario_result @@ -170,26 +171,37 @@ def _make_request( ({"max_concurrency": 2}, None), ], ) -async def test_start_run_accepts_only_external_scenario_params( +async def test_start_run_checks_scenario_params_in_the_preparation_worker( patch_central_database: MagicMock, scenario_params: dict[str, Any], rejected: str | None ) -> None: + loop_thread = threading.current_thread() + lookup_threads: list[threading.Thread] = [] + + def get_class(name: str) -> type[Scenario]: + lookup_threads.append(threading.current_thread()) + return Scam + registry = MagicMock(spec=ScenarioRegistry) - registry.__contains__.return_value = True - registry.get_class.return_value = Scam + registry.get_class.side_effect = get_class service = ScenarioRunService() request = _make_request(scenario_name="airt.scam", scenario_params=scenario_params) with ( patch(f"{_REGISTRY_PATCH_BASE}.ScenarioRegistry.get_registry_singleton", return_value=registry), - patch.object(service, "_start_run_locked_async", new_callable=AsyncMock) as start, + patch.object( + service, + "_run_initializers_async", + new_callable=AsyncMock, + side_effect=RuntimeError("initializers reached"), + ) as run_initializers, + pytest.raises(ValueError if rejected else RuntimeError, match=rejected or "initializers reached"), ): - if rejected: - with pytest.raises(ValueError, match=rejected): - await service.start_run_async(request=request) - else: - await service.start_run_async(request=request) + await service.start_run_async(request=request) - assert start.await_count == (0 if rejected else 1) + assert lookup_threads + assert loop_thread not in lookup_threads + assert run_initializers.await_count == (0 if rejected else 1) + await service.close_async() def _make_db_scenario_result( @@ -653,6 +665,9 @@ def get_aggregate_tags(cls) -> set[str]: scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._technique_class = _JailbreakTechnique + mock_all_registries["scenario_registry"].get_class.return_value.supported_parameters.return_value = ( + Jailbreak.supported_parameters() + ) objective_target = mock_all_registries["target_registry"].instances.get.return_value scenario_params = {"num_jailbreaks": 2, "num_jailbreak_attempts": 1} @@ -695,6 +710,7 @@ def get_aggregate_tags(cls) -> set[str]: service = ScenarioRunService() mock_sr = mock_all_registries["scenario_registry"] + mock_sr.get_class.return_value.supported_parameters.return_value = Jailbreak.supported_parameters() mock_memory = mock_all_registries["memory"] mock_all_registries["scenario_instance"]._technique_class = _JailbreakTechnique records: dict[str, MagicMock] = {} diff --git a/tests/unit/backend/test_scorer_service.py b/tests/unit/backend/test_scorer_service.py index b3f6d9dff2..9d73c5425b 100644 --- a/tests/unit/backend/test_scorer_service.py +++ b/tests/unit/backend/test_scorer_service.py @@ -5,7 +5,7 @@ import asyncio import threading -from collections.abc import Iterator +from collections.abc import Callable, Iterator from unittest.mock import MagicMock, patch import pytest @@ -42,6 +42,14 @@ def get_scorer_metrics(self): return None +class _ObjectConfiguredScorer(_ServiceScorer): + """A scorer with one external parameter and one that takes a Python object.""" + + def __init__(self, *, label: str = "service", tokenizer: Callable[[str], list[str]] | None = None) -> None: + super().__init__(label=label) + self.tokenizer = tokenizer + + @pytest.fixture(autouse=True) def reset_scorer_registry(patch_central_database: MagicMock) -> Iterator[None]: get_scorer_service.cache_clear() @@ -159,3 +167,32 @@ async def test_close_services_clears_scorer_cache_before_registry_replacement() await new_service.create_scorer_async(request=CreateScorerRequest(name="created", type="_ServiceScorer")) assert new_registry.instances.get("created") is not None assert old_registry.instances.get("created") is None + + +async def test_create_scorer_accepts_only_external_inputs() -> None: + ScorerRegistry.get_registry_singleton().register_class(_ObjectConfiguredScorer) + service = ScorerService() + + with pytest.raises(ValueError, match="'tokenizer' of '_ObjectConfiguredScorer' cannot be set through the API"): + await service.create_scorer_async( + request=CreateScorerRequest(name="tokenized", type="_ObjectConfiguredScorer", params={"tokenizer": "x"}) + ) + created = await service.create_scorer_async( + request=CreateScorerRequest(name="labeled", type="_ObjectConfiguredScorer", params={"label": "custom"}) + ) + + assert created.scorer_registry_name == "labeled" + assert ScorerRegistry.get_registry_singleton().instances.get("tokenized") is None + + +async def test_types_list_only_external_inputs() -> None: + registry = ScorerRegistry.get_registry_singleton() + registry.register_class(_ObjectConfiguredScorer) + + types = await ScorerService().list_scorer_types_async() + + entry = next(item for item in types.items if item.scorer_type == "_ObjectConfiguredScorer") + assert [parameter.name for parameter in entry.parameters] == ["label"] + metadata = registry.get_registered_class_metadata("_ObjectConfiguredScorer") + assert metadata is not None + assert {parameter.name for parameter in metadata.parameters} == {"label", "tokenizer"} diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index 08509531c0..178d57cc86 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -297,37 +297,37 @@ async def test_types_include_structured_parameters(self) -> None: assert weights_parameter.is_list is True assert weights_parameter.required is False - async def test_types_expose_external_inputs_in_registry_order_without_mutating_metadata(self) -> None: + async def test_types_keep_registry_parameter_order_without_mutating_metadata(self) -> None: service = TargetService() result = await service.list_target_types_async() metadata_by_name = { metadata.class_name: metadata for metadata in service._registry.get_all_registered_class_metadata() } - expected = { - name - for name, metadata in metadata_by_name.items() - if all(parameter.is_external_input for parameter in metadata.parameters if parameter.required) - } - assert {entry.target_type for entry in result.items} == expected for entry in result.items: - registry_parameters = metadata_by_name[entry.target_type].parameters - assert [parameter.name for parameter in entry.parameters] == [ - parameter.name for parameter in registry_parameters if parameter.is_external_input - ] + registry_names = iter(parameter.name for parameter in metadata_by_name[entry.target_type].parameters) + assert all(parameter.name in registry_names for parameter in entry.parameters) registry_openai = {parameter.name: parameter for parameter in metadata_by_name["OpenAIChatTarget"].parameters} assert registry_openai["endpoint"].required is False assert registry_openai["model_name"].required is False - async def test_types_keep_string_api_key_and_omit_object_only_targets(self) -> None: + async def test_types_expose_component_external_inputs(self) -> None: service = TargetService() result = await service.list_target_types_async() - entries = {entry.target_type: entry for entry in result.items} - assert "api_key" in {parameter.name for parameter in entries["OpenAIChatTarget"].parameters} - assert "custom_configuration" not in {parameter.name for parameter in entries["OpenAIChatTarget"].parameters} - assert not {"PlaywrightTarget", "PlaywrightCopilotTarget", "WebsocketTarget"} & set(entries) + parameters = {entry.target_type: {parameter.name for parameter in entry.parameters} for entry in result.items} + assert {"PlaywrightTarget", "PlaywrightCopilotTarget", "WebsocketTarget"}.isdisjoint(parameters) + assert "api_key" in parameters["OpenAIChatTarget"] + assert "custom_configuration" not in parameters["OpenAIChatTarget"] + assert "n_seconds" in parameters["OpenAIVideoTarget"] + + async def test_registry_metadata_keeps_parameters_the_api_cannot_set(self) -> None: + metadata = TargetService()._registry.get_registered_class_metadata("OpenAIChatTarget") + assert metadata is not None + registry_parameters = {parameter.name: parameter for parameter in metadata.parameters} + + assert not registry_parameters["custom_configuration"].is_external_input async def test_types_cold_and_warm_results_are_equal(self) -> None: service = TargetService() @@ -409,6 +409,23 @@ async def test_create_target_rejects_parameters_that_take_python_objects(self, s assert service.get_target_object(target_registry_name="text") is None + async def test_create_target_accepts_scalar_alternative_of_union(self, sqlite_instance) -> None: + service = TargetService() + request = CreateTargetRequest( + name="video", + type="OpenAIVideoTarget", + params={ + "endpoint": "https://example.openai.azure.com/openai/v1", + "api_key": "test-key", + "model_name": "sora-2", + "n_seconds": 8, + }, + ) + + await service.create_target_async(request=request) + + assert service.get_target_object(target_registry_name="video")._n_seconds == "8" + async def test_create_target_raises_for_invalid_type(self) -> None: """Test that create_target raises for invalid target type.""" service = TargetService() diff --git a/tests/unit/models/test_parameter.py b/tests/unit/models/test_parameter.py index d8f055251f..0553f0e68d 100644 --- a/tests/unit/models/test_parameter.py +++ b/tests/unit/models/test_parameter.py @@ -3,7 +3,7 @@ """Unit tests for the unified Parameter model and its coercion methods.""" -from collections.abc import Callable, Sequence +from collections.abc import Callable, Collection, Sequence from enum import Enum from pathlib import Path from typing import Any, Literal, Union @@ -557,9 +557,17 @@ class TestIsExternalInput: _Speed, list[str], list[_Speed] | None, + Sequence[str], + Collection[str] | None, + Sequence[_Speed], str | list[str], str | Callable[[], str] | None, _Speed | str, + int | tuple[int, int], + int | Literal["4", "8"], + str | dict[str, str], + str | _Unsupported, + Collection[str] | _Unsupported, ], ) def test_supported_external_types(self, param_type: object) -> None: @@ -575,17 +583,18 @@ def test_supported_external_types(self, param_type: object) -> None: Callable[[], str], tuple[int, int], dict[str, str], - Sequence[str], + set[str], list[Path], list[Path | str], - int | tuple[int, int], + Collection[Path], + Sequence[Path | str] | None, _Speed | Path, str | Path | int, - str | dict[str, str], - str | _Unsupported, str | list[Path], + _Unsupported | Path, list[list[str]], - int | Literal["4", "8"], + Sequence[list[str]], + tuple[int, int] | _Unsupported, ], ) def test_other_types_take_python_objects_only(self, param_type: object) -> None: From 0628f5707fcba1cd25c83ee5bda3193129b36334 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 19:18:56 +0000 Subject: [PATCH 5/8] Note that path scalars are external inputs and simplify two tests --- pyrit/models/parameter.py | 7 ++++--- tests/unit/backend/test_scenario_run_service.py | 5 ++--- tests/unit/backend/test_target_service.py | 5 +++-- 3 files changed, 9 insertions(+), 8 deletions(-) diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index 0bf841d7ca..0f66c446cf 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -283,9 +283,10 @@ def is_external_input(self) -> bool: """ Whether REST, CLI, and GUI callers may supply this parameter. - True for registry references, declared structured inputs, scalars, flat - ``list`` / ``Collection`` / ``Sequence`` of non-path scalars, and unions with one - of those as an alternative and no path alternative. External callers supply that + True for registry references, declared structured inputs, scalars (``Path`` and + ``Path | str`` included), flat ``list`` / ``Collection`` / ``Sequence`` of non-path + scalars, and other unions with one of those as an alternative and no path + alternative. External callers supply that alternative, as for ``api_key: str | Callable[...]`` or ``font_size: int | tuple[int, int]``; the other alternatives are for in-process callers. Other parameters take Python objects from in-process callers only. diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index 1134f4cea1..6359cefc25 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -665,9 +665,8 @@ def get_aggregate_tags(cls) -> set[str]: scenario_instance = mock_all_registries["scenario_instance"] scenario_instance._technique_class = _JailbreakTechnique - mock_all_registries["scenario_registry"].get_class.return_value.supported_parameters.return_value = ( - Jailbreak.supported_parameters() - ) + mock_sr = mock_all_registries["scenario_registry"] + mock_sr.get_class.return_value.supported_parameters.return_value = Jailbreak.supported_parameters() objective_target = mock_all_registries["target_registry"].instances.get.return_value scenario_params = {"num_jailbreaks": 2, "num_jailbreak_attempts": 1} diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index 178d57cc86..7f822d88c0 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -305,8 +305,9 @@ async def test_types_keep_registry_parameter_order_without_mutating_metadata(sel } for entry in result.items: - registry_names = iter(parameter.name for parameter in metadata_by_name[entry.target_type].parameters) - assert all(parameter.name in registry_names for parameter in entry.parameters) + entry_names = [parameter.name for parameter in entry.parameters] + registry_names = [parameter.name for parameter in metadata_by_name[entry.target_type].parameters] + assert entry_names == [name for name in registry_names if name in entry_names] registry_openai = {parameter.name: parameter for parameter in metadata_by_name["OpenAIChatTarget"].parameters} assert registry_openai["endpoint"].required is False From ffacb803c481a5f894d4c7d6579928300a034d53 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 20:12:30 +0000 Subject: [PATCH 6/8] Describe what the type catalogs return to external callers --- doc/gui/0_gui.md | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/doc/gui/0_gui.md b/doc/gui/0_gui.md index 4afd1e8d35..a440c0a8c2 100644 --- a/doc/gui/0_gui.md +++ b/doc/gui/0_gui.md @@ -588,9 +588,10 @@ at `GET /api/runtime`. ## Registry API Migration Notes Use `/api/converters/types` and `/api/targets/types` for registry build metadata. -These endpoints return all constructor parameters from the registry, including -lists, unions, and component references. The temporary `/catalog` routes retain -their scalar-only filtering for the current UI. +These endpoints return the constructor parameters external callers can set, +including flat lists, unions with a supported alternative, and component +references, and leave out types external callers can't create. Registry metadata +keeps every parameter. Create requests should supply an explicit registry `name`. Converter creation returns the complete `ConverterInstance`; read its type from `identifier.class_name`, not the old top-level `converter_type` field. Treat From 3ff7c7466a6984f3358665357179e34bad9f12c7 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 20:37:21 +0000 Subject: [PATCH 7/8] Remove stale catalog route note from the GUI docs --- doc/gui/0_gui.md | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/doc/gui/0_gui.md b/doc/gui/0_gui.md index a440c0a8c2..f59001f3e9 100644 --- a/doc/gui/0_gui.md +++ b/doc/gui/0_gui.md @@ -608,10 +608,8 @@ allowlisted image, audio, and video extensions inline. Other files, including PD SVG, HTML, text, and executables, download as `application/octet-stream` attachments. **Temporary compatibility, scheduled for removal with the chat migration:** -the `/api/converters/catalog` and `/api/targets/catalog` routes project the same -registry metadata for the current UI. Create requests without a name receive a -generated `compat_...` name. New clients should not depend on these routes or -unnamed creation. +target create requests without a name receive a generated `compat_...` name. +New clients should supply an explicit name. ## Connection Health From 0c0c7f04218a0915be2f7db6131878a15c67ef04 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 21:18:18 +0000 Subject: [PATCH 8/8] Let the preparation callback run before closing in the worker test --- tests/unit/backend/test_scenario_run_service.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index 6359cefc25..9f687b7808 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -201,6 +201,8 @@ def get_class(name: str) -> type[Scenario]: assert lookup_threads assert loop_thread not in lookup_threads assert run_initializers.await_count == (0 if rejected else 1) + # The finished preparation's done callback can still be queued when its error reaches the caller. + await asyncio.sleep(0) await service.close_async()