diff --git a/pyrit/backend/mappers/converter_mappers.py b/pyrit/backend/mappers/converter_mappers.py index 7c7d005e34..b1c11b2d7d 100644 --- a/pyrit/backend/mappers/converter_mappers.py +++ b/pyrit/backend/mappers/converter_mappers.py @@ -15,7 +15,13 @@ from pyrit.models import ConverterIdentifier -def converter_object_to_instance(converter_id: str, converter_obj: Converter) -> ConverterInstance: +def converter_object_to_instance( + *, + converter_id: str, + converter_obj: Converter, + is_llm_based: bool, + description: str | None, +) -> ConverterInstance: """ Build a ConverterInstance DTO from a registry converter object. @@ -24,8 +30,10 @@ def converter_object_to_instance(converter_id: str, converter_obj: Converter) -> on the wire. Args: - converter_id: The unique converter instance identifier. - converter_obj: The domain Converter object from the registry. + converter_id (str): The unique converter instance identifier. + converter_obj (Converter): The domain Converter object from the registry. + is_llm_based (bool): Whether the converter class requires an LLM target. + description (str | None): The converter class description. Returns: ConverterInstance DTO wrapping the converter's identifier. @@ -33,4 +41,6 @@ def converter_object_to_instance(converter_id: str, converter_obj: Converter) -> return ConverterInstance( converter_id=converter_id, identifier=ConverterIdentifier.from_component_identifier(converter_obj.get_identifier()), + is_llm_based=is_llm_based, + description=description, ) diff --git a/pyrit/backend/models/__init__.py b/pyrit/backend/models/__init__.py index 6307bb1733..dd142e9b82 100644 --- a/pyrit/backend/models/__init__.py +++ b/pyrit/backend/models/__init__.py @@ -51,6 +51,8 @@ ConverterInstanceListResponse, ConverterPreviewRequest, ConverterPreviewResponse, + ConverterTypeEntry, + ConverterTypeResponse, CreateConverterRequest, CreateConverterResponse, PreviewStep, @@ -97,6 +99,8 @@ "ConverterInstanceListResponse": "pyrit.backend.models.converters", "ConverterPreviewRequest": "pyrit.backend.models.converters", "ConverterPreviewResponse": "pyrit.backend.models.converters", + "ConverterTypeEntry": "pyrit.backend.models.converters", + "ConverterTypeResponse": "pyrit.backend.models.converters", "CreateConverterRequest": "pyrit.backend.models.converters", "CreateConverterResponse": "pyrit.backend.models.converters", "PreviewStep": "pyrit.backend.models.converters", diff --git a/pyrit/backend/models/common.py b/pyrit/backend/models/common.py index 36767467cc..33e751dbc7 100644 --- a/pyrit/backend/models/common.py +++ b/pyrit/backend/models/common.py @@ -11,6 +11,8 @@ from pydantic import BaseModel, Field +REGISTRY_INSTANCE_NAME_PATTERN = r"^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$" + class PaginationInfo(BaseModel): """Pagination metadata for list responses.""" diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index c7dac9493f..9f165465f2 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -11,6 +11,7 @@ from pydantic import BaseModel, Field +from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN from pyrit.models import ConverterIdentifier, Parameter, PromptDataType __all__ = [ @@ -18,6 +19,8 @@ "ConverterCatalogResponse", "ConverterInstance", "ConverterInstanceListResponse", + "ConverterTypeEntry", + "ConverterTypeResponse", "CreateConverterRequest", "CreateConverterResponse", "ConverterPreviewRequest", @@ -27,11 +30,11 @@ # ============================================================================ -# Converter Catalog (Available Types) +# Converter Types # ============================================================================ -class ConverterCatalogEntry(BaseModel): +class ConverterTypeEntry(BaseModel): """A converter type available from the backend registry.""" converter_type: str = Field(..., description="Converter class name (e.g., 'Base64Converter')") @@ -48,10 +51,17 @@ class ConverterCatalogEntry(BaseModel): description: str | None = Field(None, description="Short description of the converter from its docstring") -class ConverterCatalogResponse(BaseModel): +class ConverterTypeResponse(BaseModel): """Response for listing available converter types from the registry.""" - items: list[ConverterCatalogEntry] = Field(..., description="List of available converter types") + items: list[ConverterTypeEntry] = Field(..., description="List of available converter types") + + +# LEGACY COMPATIBILITY: ``Catalog`` is the pre-registry name for ``Type``. These +# aliases exist only so the un-migrated chat UI keeps working; delete them with the +# /catalog route when that UI switches to the /types API. +ConverterCatalogEntry = ConverterTypeEntry +ConverterCatalogResponse = ConverterTypeResponse # ============================================================================ @@ -68,8 +78,10 @@ class ConverterInstance(BaseModel): for the converter's class, supported data types, and constructor params. """ - converter_id: str = Field(..., description="Unique converter instance identifier") + converter_id: str = Field(..., description="Converter instance registry name") identifier: ConverterIdentifier = Field(..., description="The converter's identity/configuration projection") + 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 type") class ConverterInstanceListResponse(BaseModel): @@ -81,7 +93,17 @@ class ConverterInstanceListResponse(BaseModel): class CreateConverterRequest(BaseModel): """Request to create a new converter instance.""" + # LEGACY COMPATIBILITY: The current chat UI does not send a name. Make this + # field required when the chat-migration stack layer sends explicit names. + name: str | None = Field( + None, + min_length=1, + pattern=REGISTRY_INSTANCE_NAME_PATTERN, + description="Unique registry name; omitted only for legacy chat compatibility", + ) type: str = Field(..., description="Converter type (e.g., 'Base64Converter')") + # LEGACY COMPATIBILITY: The former create response echoed this field. Remove + # it after clients use the complete ConverterInstance response. display_name: str | None = Field(None, description="Human-readable display name") params: dict[str, Any] = Field( default_factory=dict, @@ -90,7 +112,12 @@ class CreateConverterRequest(BaseModel): class CreateConverterResponse(BaseModel): - """Response after creating a converter instance.""" + """ + Legacy response model for downstream imports. + + POST /converters now returns ``ConverterInstance``. Remove this model when + downstream clients no longer import the former response type. + """ converter_id: str = Field(..., description="Unique converter instance identifier") converter_type: str = Field(..., description="Converter class name") diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index 87797d32d1..10bcd14edc 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -12,7 +12,7 @@ from pydantic import BaseModel, Field -from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN, PaginationInfo from pyrit.models import JSONValue, Parameter from pyrit.models.catalog.target import TargetInstance @@ -21,6 +21,8 @@ "TargetCatalogEntry", "TargetCatalogResponse", "TargetListResponse", + "TargetTypeEntry", + "TargetTypeResponse", ] @@ -28,7 +30,7 @@ def _default_auth_modes() -> list[Literal["api_key", "identity"]]: return ["api_key"] -class TargetCatalogEntry(BaseModel): +class TargetTypeEntry(BaseModel): """A target type available from the backend registry.""" target_type: str = Field(..., description="Target class name (e.g., 'OpenAIChatTarget')") @@ -43,10 +45,17 @@ class TargetCatalogEntry(BaseModel): description: str | None = Field(None, description="Short description of the target from its docstring") -class TargetCatalogResponse(BaseModel): +class TargetTypeResponse(BaseModel): """Response for listing available target types from the registry.""" - items: list[TargetCatalogEntry] = Field(..., description="List of available target types") + items: list[TargetTypeEntry] = Field(..., description="List of available target types") + + +# LEGACY COMPATIBILITY: ``Catalog`` is the pre-registry name for ``Type``. These +# aliases exist only so the un-migrated configuration UI keeps working; delete them +# with the /catalog route when that UI switches to the /types API. +TargetCatalogEntry = TargetTypeEntry +TargetCatalogResponse = TargetTypeResponse class TargetListResponse(BaseModel): @@ -59,6 +68,14 @@ class TargetListResponse(BaseModel): class CreateTargetRequest(BaseModel): """Request to create a new target instance.""" + # LEGACY COMPATIBILITY: The current target configuration UI does not send a + # name. Make this field required after that UI sends explicit registry names. + name: str | None = Field( + None, + min_length=1, + pattern=REGISTRY_INSTANCE_NAME_PATTERN, + description="Unique registry name; omitted only for legacy UI compatibility", + ) type: str = Field(..., description="Target type (e.g., 'OpenAIChatTarget')") params: dict[str, JSONValue] = Field(default_factory=dict, description="Target constructor parameters") auth_mode: Literal["api_key", "identity"] = Field( diff --git a/pyrit/backend/routes/converters.py b/pyrit/backend/routes/converters.py index c741353919..aad4c9ac39 100644 --- a/pyrit/backend/routes/converters.py +++ b/pyrit/backend/routes/converters.py @@ -17,8 +17,8 @@ ConverterInstanceListResponse, ConverterPreviewRequest, ConverterPreviewResponse, + ConverterTypeResponse, CreateConverterRequest, - CreateConverterResponse, ) from pyrit.backend.services.converter_service import get_converter_service @@ -42,16 +42,35 @@ async def list_converters() -> ConverterInstanceListResponse: # pyrit-async-suf return await service.list_converters_async() +@router.get( + "/types", + response_model=ConverterTypeResponse, +) +async def list_converter_types() -> ConverterTypeResponse: # pyrit-async-suffix-exempt + """ + List converter types projected from ``ConverterRegistry`` metadata. + + Returns: + ConverterTypeResponse: Available converter types and build parameters. + """ + service = get_converter_service() + return await service.list_converter_types_async() + + @router.get( "/catalog", response_model=ConverterCatalogResponse, ) async def list_converter_catalog() -> ConverterCatalogResponse: # pyrit-async-suffix-exempt """ - List all available converter types from the backend converter registry. + Return the legacy catalog projection used by the current chat UI. + + LEGACY COMPATIBILITY: pre-registry alias for ``/converters/types`` that hides + registry-reference parameters. Deleted with the rest of the ``catalog`` concept + when the chat-migration layer of this stack switches to ``/converters/types``. Returns: - ConverterCatalogResponse: List of available converter types. + ConverterCatalogResponse: The scalar-only legacy catalog projection. """ service = get_converter_service() return await service.list_converter_catalog_async() @@ -59,13 +78,13 @@ async def list_converter_catalog() -> ConverterCatalogResponse: # pyrit-async-s @router.post( "", - response_model=CreateConverterResponse, + response_model=ConverterInstance, status_code=status.HTTP_201_CREATED, responses={ 400: {"model": ProblemDetail, "description": "Invalid converter type or parameters"}, }, ) -async def create_converter(request: CreateConverterRequest) -> CreateConverterResponse: # pyrit-async-suffix-exempt +async def create_converter(request: CreateConverterRequest) -> ConverterInstance: # pyrit-async-suffix-exempt """ Create a new converter instance. @@ -73,7 +92,7 @@ async def create_converter(request: CreateConverterRequest) -> CreateConverterRe Supports nested converters via converter_id references in params. Returns: - CreateConverterResponse: The created converter instance details. + ConverterInstance: The created converter instance details. """ service = get_converter_service() @@ -117,6 +136,23 @@ async def get_converter(converter_id: str) -> ConverterInstance: # pyrit-async- return converter +@router.delete( + "/{converter_id}", + status_code=status.HTTP_204_NO_CONTENT, + responses={ + 404: {"model": ProblemDetail, "description": "Converter not found"}, + }, +) +async def delete_converter(converter_id: str) -> None: # pyrit-async-suffix-exempt + """Delete a converter instance by registry name.""" + service = get_converter_service() + if not await service.delete_converter_async(converter_id=converter_id): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Converter '{converter_id}' not found", + ) + + @router.post( "/preview", response_model=ConverterPreviewResponse, diff --git a/pyrit/backend/routes/media.py b/pyrit/backend/routes/media.py index e8969bbfeb..502b55f4f1 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -8,6 +8,15 @@ so the frontend can reference them by URL instead of requiring inline base64 data URIs. For Azure deployments, media is served directly from Azure Blob Storage via signed URLs and this endpoint is not used. + +This route is the only place PyRIT hands stored bytes to a browser, so it is the +only place that restricts content. Storage stays unrestricted on purpose: any +file type is a legitimate attack payload (uploading an ``.html`` file so an attack +can push it to a blob target is a valid operation). Files therefore keep their +real name and extension on disk, and this route never renames them -- anything +reading a stored path gets the original file. Only the HTTP *response* is +adjusted: a document type a browser would execute in this origin is returned as an +opaque download instead of a rendered page. """ import logging @@ -26,6 +35,13 @@ # Only serve files from known media subdirectories under results_path. _ALLOWED_SUBDIRECTORIES = {"prompt-memory-entries", "seed-prompt-entries"} +# Types a browser executes in this origin. They are still stored and still served, +# but always as an opaque download so stored content cannot script against the UI. +_ACTIVE_DOCUMENT_EXTENSIONS = {".htm", ".html", ".svg", ".xhtml", ".xml"} + +# Types the browser is asked to download rather than render inline. +_ATTACHMENT_EXTENSIONS = {".csv", ".md", ".pdf", ".txt"} | _ACTIVE_DOCUMENT_EXTENSIONS + # Only serve known media file types (allowlist approach). _ALLOWED_EXTENSIONS = { # Images @@ -35,7 +51,6 @@ ".gif", ".bmp", ".webp", - ".svg", ".ico", ".tiff", # Audio @@ -56,8 +71,7 @@ ".md", ".csv", ".pdf", - ".html", -} +} | _ACTIVE_DOCUMENT_EXTENSIONS def _validate_media_path(*, path: str, allowed_root: Path) -> Path: @@ -110,6 +124,12 @@ async def serve_media_async( configured results directory (e.g. ``dbdata/prompt-memory-entries/``) to prevent path traversal attacks and exfiltration of sensitive files. + The stored file is never modified or renamed. Active document types + (see ``_ACTIVE_DOCUMENT_EXTENSIONS``) are returned as opaque downloads so the + browser does not execute them in this origin; the bytes and the file name are + unchanged, so a caller that needs the real file (e.g. to attach an ``.html`` + payload to a target) reads it from its stored path. + Args: path: Absolute path to the file. @@ -134,8 +154,16 @@ async def serve_media_async( if not validated_path.is_file(): raise HTTPException(status_code=404, detail="File not found.") - mime_type, _ = mimetypes.guess_type(validated_path) + extension = validated_path.suffix.lower() + if extension in _ACTIVE_DOCUMENT_EXTENSIONS: + media_type = "application/octet-stream" + else: + guessed_type, _ = mimetypes.guess_type(validated_path) + media_type = guessed_type or "application/octet-stream" return FileResponse( path=validated_path, - media_type=mime_type or "application/octet-stream", + media_type=media_type, + filename=validated_path.name if extension in _ATTACHMENT_EXTENSIONS else None, + content_disposition_type="attachment", + headers={"X-Content-Type-Options": "nosniff"}, ) diff --git a/pyrit/backend/routes/targets.py b/pyrit/backend/routes/targets.py index 3bac8a23b7..0e2e35ed30 100644 --- a/pyrit/backend/routes/targets.py +++ b/pyrit/backend/routes/targets.py @@ -15,6 +15,7 @@ CreateTargetRequest, TargetCatalogResponse, TargetListResponse, + TargetTypeResponse, ) from pyrit.backend.services.target_service import get_target_service from pyrit.models.catalog.target import TargetInstance @@ -45,6 +46,24 @@ async def list_targets( # pyrit-async-suffix-exempt return await service.list_targets_async(limit=limit, cursor=cursor) +@router.get( + "/types", + response_model=TargetTypeResponse, + responses={ + 500: {"model": ProblemDetail, "description": "Internal server error"}, + }, +) +async def list_target_types() -> TargetTypeResponse: # pyrit-async-suffix-exempt + """ + List target types projected from ``TargetRegistry`` metadata. + + Returns: + TargetTypeResponse: Available target types and build parameters. + """ + service = get_target_service() + return await service.list_target_types_async() + + @router.get( "/catalog", response_model=TargetCatalogResponse, @@ -54,10 +73,14 @@ async def list_targets( # pyrit-async-suffix-exempt ) async def list_target_catalog() -> TargetCatalogResponse: # pyrit-async-suffix-exempt """ - List all available target types from the backend target registry. + Return the legacy catalog projection used by the current configuration UI. + + LEGACY COMPATIBILITY: pre-registry alias for ``/targets/types`` that hides + registry-reference parameters. Deleted with the rest of the ``catalog`` concept + when the configuration UI switches to ``/targets/types``. Returns: - TargetCatalogResponse: List of available target types. + TargetCatalogResponse: The scalar-only legacy catalog projection. """ service = get_target_service() return await service.list_target_catalog_async() diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index 04fc518f54..520a926b42 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -16,24 +16,29 @@ import binascii import mimetypes import uuid +from contextlib import suppress from functools import lru_cache from pathlib import Path from typing import TYPE_CHECKING, Any from urllib.parse import parse_qs, urlparse +import aiofiles +import aiofiles.os + from pyrit.backend.mappers.converter_mappers import converter_object_to_instance from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS from pyrit.backend.models.converters import ( - ConverterCatalogEntry, ConverterCatalogResponse, ConverterInstance, ConverterInstanceListResponse, ConverterPreviewRequest, ConverterPreviewResponse, + ConverterTypeEntry, + ConverterTypeResponse, CreateConverterRequest, - CreateConverterResponse, PreviewStep, ) +from pyrit.common.path import DB_DATA_PATH from pyrit.memory import data_serializer_factory from pyrit.models import PromptDataType from pyrit.registry.components import ConverterRegistry @@ -42,6 +47,11 @@ from pyrit.converter import ConverterResult +_OWNED_ARTIFACT_PATHS_KEY = "owned_artifact_paths" +_REGISTRY_UPLOAD_DIRECTORY = DB_DATA_PATH / "registry-uploads" +_DEFAULT_UPLOAD_EXTENSION = ".bin" + + class ConverterService: """ Service for managing converter instances. @@ -63,7 +73,14 @@ def _build_instance_from_object(self, *, converter_id: str, converter_obj: Any) Returns: ConverterInstance with metadata derived from the object's identifier. """ - return converter_object_to_instance(converter_id, converter_obj) + metadata = self._registry.get_registered_class_metadata(converter_obj.__class__.__name__) + description = metadata.class_description or None if metadata else None + return converter_object_to_instance( + converter_id=converter_id, + converter_obj=converter_obj, + is_llm_based=metadata.is_llm_based if metadata else False, + description=description, + ) # ======================================================================== # Public API Methods @@ -82,7 +99,7 @@ async def list_converters_async(self) -> ConverterInstanceListResponse: ] return ConverterInstanceListResponse(items=items) - async def list_converter_catalog_async(self) -> ConverterCatalogResponse: + async def list_converter_types_async(self) -> ConverterTypeResponse: """ List all available converter types from the converter class registry. @@ -91,20 +108,42 @@ async def list_converter_catalog_async(self) -> ConverterCatalogResponse: frontend), not this service. Returns: - ConverterCatalogResponse containing all available converter classes. + ConverterTypeResponse containing all available converter classes. """ - items: list[ConverterCatalogEntry] = [ - ConverterCatalogEntry( + items: list[ConverterTypeEntry] = [ + ConverterTypeEntry( converter_type=metadata.class_name, supported_input_types=list(metadata.supported_input_types), supported_output_types=list(metadata.supported_output_types), - parameters=[p for p in metadata.parameters if p.is_string_coercible], + parameters=[p for p in metadata.parameters if p.is_string_coercible or p.reference is not None], is_llm_based=metadata.is_llm_based, description=metadata.class_description or None, ) for metadata in self._registry.get_all_registered_class_metadata() ] + return ConverterTypeResponse(items=items) + + async def list_converter_catalog_async(self) -> ConverterCatalogResponse: + """ + Return the legacy projection used by the current chat UI. + + LEGACY COMPATIBILITY: ``catalog`` is the pre-registry name for ``types``, and + the whole concept goes away -- there is no ``ConverterCatalog`` class and + nothing new should use this. It differs from ``list_converter_types_async`` in + exactly one way: it drops registry-reference parameters, which the un-migrated + chat UI cannot render. Delete this method, the ``/catalog`` route, and the + ``ConverterCatalog*`` aliases together when the chat-migration layer of this + stack switches to ``/converters/types``. + + Returns: + ConverterCatalogResponse: The scalar-only legacy projection. + """ + types_response = await self.list_converter_types_async() + items = [ + entry.model_copy(update={"parameters": [p for p in entry.parameters if p.is_string_coercible]}) + for entry in types_response.items + ] return ConverterCatalogResponse(items=items) async def get_converter_async(self, *, converter_id: str) -> ConverterInstance | None: @@ -128,7 +167,22 @@ def get_converter_object(self, *, converter_id: str) -> Any | None: """ return self._registry.instances.get(converter_id) - async def create_converter_async(self, *, request: CreateConverterRequest) -> CreateConverterResponse: + async def delete_converter_async(self, *, converter_id: str) -> bool: + """ + Delete a converter instance by registry name. + + Returns: + bool: True when an instance was removed, otherwise False. + """ + entry = self._registry.instances.get_entry(converter_id) + if entry is None: + return False + + owned_paths = self._get_owned_artifact_paths(entry.metadata) + await self._remove_owned_artifacts_async(paths=owned_paths) + return self._registry.instances.unregister(converter_id) is not None + + async def create_converter_async(self, *, request: CreateConverterRequest) -> ConverterInstance: """ Create a new converter instance from API request. @@ -139,26 +193,36 @@ async def create_converter_async(self, *, request: CreateConverterRequest) -> Cr request: The create converter request with type and params. Returns: - CreateConverterResponse with the new converter's details. + ConverterInstance with the new converter's details. Raises: - ValueError: If the converter type is not found. + ValueError: If the converter type is not found or the registry name is + unavailable. """ - converter_id = str(uuid.uuid4()) - - # Persist data-URI params to disk (frontend concern), then delegate - # construction (incl. param coercion and reference resolution) to the - # converter registry. if request.type not in self._registry: raise ValueError(f"Converter type '{request.type}' not found") - params = await self._persist_data_uri_params_async(converter_type=request.type, params=request.params) - converter_obj = self._registry.create_instance(request.type, **params) - self._registry.instances.register(converter_obj, name=converter_id) + # LEGACY COMPATIBILITY: The current chat UI omits the name. Remove this + # generated fallback when that UI sends an explicit registry name. + converter_id = request.name or f"compat_{uuid.uuid4().hex}" + self._registry.instances.validate_name_available(converter_id) + params, owned_paths = await self._persist_data_uri_params_async( + converter_type=request.type, + params=request.params, + ) + try: + converter_obj = self._registry.create_named_instance( + name=converter_id, + converter_type=request.type, + registry_metadata={_OWNED_ARTIFACT_PATHS_KEY: [str(path) for path in owned_paths]}, + **params, + ) + except Exception: + await self._remove_owned_artifacts_async(paths=owned_paths) + raise - return CreateConverterResponse( + return self._build_instance_from_object( converter_id=converter_id, - converter_type=request.type, - display_name=request.display_name, + converter_obj=converter_obj, ) async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> ConverterPreviewResponse: @@ -257,15 +321,16 @@ async def _persist_data_uri_params_async( *, converter_type: str, params: dict[str, Any], - ) -> dict[str, Any]: + ) -> tuple[dict[str, Any], list[Path]]: """ - Persist data-URI parameter values to disk. + Persist uploaded ``Path`` parameter values to managed local storage. The frontend file picker sends file contents as data URIs - (e.g. ``data:image/png;base64,...``). Constructor parameters typed as - ``Path`` or ``str`` params whose names suggest a file path receive the - decoded file persisted to the results store, with the value replaced - by the resulting file path. + (e.g. ``data:image/png;base64,...``). A constructor parameter typed as ``Path`` + is therefore an *upload*: the decoded file is written to a local working + directory this service owns, and the client never names a server path. Every + ``Path`` parameter is handled the same way, so a converter opts in simply by + declaring the type; there is no per-converter or per-parameter table. The set of constructor parameters (and their types) is sourced from the registry's derived ``Parameter`` metadata rather than re-introspecting the @@ -276,45 +341,108 @@ async def _persist_data_uri_params_async( params (dict[str, Any]): The raw constructor params from the request. Returns: - dict[str, Any]: Params dict with data-URI values replaced by file paths. + tuple[dict[str, Any], list[Path]]: Updated parameters and the explicit + set of request-created files owned by the future registry entry. + + Raises: + ValueError: If a ``Path`` value is not a valid data URI. """ metadata = self._registry.get_registered_class_metadata(converter_type) param_types = {p.name: p.param_type for p in metadata.parameters} if metadata else {} result = dict(params) - for name, value in result.items(): - if not isinstance(value, str) or not value.startswith("data:"): - continue - if name not in param_types: - continue - - # Parse data URI: data:[][;base64], - header, _, payload = value.partition(",") - if not payload: - continue - - # Derive extension from the MIME type in the header - mime_type = header.split(":")[1].split(";")[0] if ":" in header else "" - ext = mimetypes.guess_extension(mime_type, strict=False) if mime_type else None - if not ext: - ext = ".bin" - - serializer = data_serializer_factory( - category="prompt-memory-entries", - data_type="binary_path", - extension=ext, - ) - await serializer.save_data_async(data=base64.b64decode(payload)) - file_path = str(serializer.value) - - # The registry already unwraps Optional, so ``param_type`` is ``Path`` - # for a ``Path | None`` constructor parameter. - if param_types[name] is Path: - result[name] = Path(file_path) - else: + owned_paths: list[Path] = [] + try: + for name, value in result.items(): + if param_types.get(name) is not Path: + continue + if value is None: + continue + if not isinstance(value, str) or not value.startswith("data:"): + raise ValueError(f"Path parameter '{name}' must be uploaded as a data URI") + + content, extension = self._decode_data_uri(parameter_name=name, data_uri=value) + file_path = await self._save_owned_artifact_async(content=content, extension=extension) + owned_paths.append(file_path) result[name] = file_path + except Exception: + await self._remove_owned_artifacts_async(paths=owned_paths) + raise + + return result, owned_paths + + @staticmethod + def _decode_data_uri(*, parameter_name: str, data_uri: str) -> tuple[bytes, str]: + """ + Decode one base64 data URI into raw content and the extension to store it under. + + Uploaded content is stored verbatim, whatever its type. PyRIT operators are + trusted and every file type is a legitimate payload: uploading an HTML file so + an attack can push it to a blob target is a valid operation. The only thing the + server decides here is the file *name*, which is generated, so a declared MIME + type can never influence where the upload lands. Restrictions on rendering + untrusted content belong to the media route that serves it back, not to storage. + + Returns: + tuple[bytes, str]: The decoded content and its file extension. + + Raises: + ValueError: If the value is not a base64 data URI or its payload is not + valid base64. + """ + header, separator, payload = data_uri.partition(",") + media_type, _, encoding = header.removeprefix("data:").partition(";") + if not separator or not payload or not header.startswith("data:") or encoding.lower() != "base64": + raise ValueError(f"Path parameter '{parameter_name}' must be a base64 data URI") + + try: + content = base64.b64decode(payload, validate=True) + except (binascii.Error, ValueError) as exc: + raise ValueError(f"Path parameter '{parameter_name}' contains invalid base64 data") from exc + + media_type = media_type.strip().lower() + extension = mimetypes.guess_extension(media_type) if media_type else None + return content, extension or _DEFAULT_UPLOAD_EXTENSION + + @staticmethod + async def _save_owned_artifact_async(*, content: bytes, extension: str) -> Path: + """ + Write one validated upload to the managed local registry directory. + + Returns: + Path: The absolute path of the new local artifact. + """ + await aiofiles.os.makedirs(_REGISTRY_UPLOAD_DIRECTORY, exist_ok=True) + file_path = (_REGISTRY_UPLOAD_DIRECTORY / f"{uuid.uuid4().hex}{extension}").resolve() + async with aiofiles.open(file_path, "xb") as file: + await file.write(content) + return file_path - return result + @staticmethod + def _get_owned_artifact_paths(metadata: dict[str, Any]) -> list[Path]: + """ + Read explicit artifact ownership from registry-entry metadata. + + Returns: + list[Path]: Paths explicitly owned by the registry entry. + """ + raw_paths = metadata.get(_OWNED_ARTIFACT_PATHS_KEY, []) + if not isinstance(raw_paths, list) or not all(isinstance(path, str) for path in raw_paths): + raise ValueError("Registry entry has invalid owned artifact metadata") + return [Path(path) for path in raw_paths] + + @staticmethod + async def _remove_owned_artifacts_async(*, paths: list[Path]) -> None: + """Remove explicitly owned files, limited to the managed upload directory.""" + allowed_root = _REGISTRY_UPLOAD_DIRECTORY.resolve() + for path in paths: + resolved_path = path.resolve() + try: + resolved_path.relative_to(allowed_root) + except ValueError as exc: + raise ValueError(f"Owned artifact path is outside the managed upload directory: {path}") from exc + with suppress(FileNotFoundError): + await aiofiles.os.remove(resolved_path) def _gather_converters(self, *, converter_ids: list[str]) -> list[tuple[str, str, Any]]: """ diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 77d078a18c..1a041d839f 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -14,6 +14,7 @@ import asyncio import logging +import uuid from functools import lru_cache from typing import Any, Literal, cast @@ -21,9 +22,10 @@ from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.targets import ( CreateTargetRequest, - TargetCatalogEntry, TargetCatalogResponse, TargetListResponse, + TargetTypeEntry, + TargetTypeResponse, ) from pyrit.models.catalog.target import TargetInstance from pyrit.registry import TargetRegistry @@ -125,7 +127,7 @@ def get_target_object(self, *, target_registry_name: str) -> Any | None: """ return self._registry.instances.get(target_registry_name) - async def list_target_catalog_async(self) -> TargetCatalogResponse: + async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. @@ -136,18 +138,40 @@ async def list_target_catalog_async(self) -> TargetCatalogResponse: not this service. Returns: - TargetCatalogResponse containing all available target classes. + TargetTypeResponse containing all available target classes. """ metadata_items = await asyncio.to_thread(self._registry.get_all_registered_class_metadata) - items: list[TargetCatalogEntry] = [ - TargetCatalogEntry( + items: list[TargetTypeEntry] = [ + TargetTypeEntry( target_type=metadata.class_name, - parameters=[p for p in metadata.parameters if p.is_string_coercible], + parameters=[p for p in metadata.parameters if p.is_string_coercible or p.reference is not None], supported_auth_modes=cast("list[Literal['api_key', 'identity']]", list(metadata.supported_auth_modes)), description=metadata.class_description or None, ) for metadata in metadata_items ] + return TargetTypeResponse(items=items) + + async def list_target_catalog_async(self) -> TargetCatalogResponse: + """ + Return the legacy projection used by the current configuration UI. + + LEGACY COMPATIBILITY: ``catalog`` is the pre-registry name for ``types``, and + the whole concept goes away -- there is no ``TargetCatalog`` class and nothing + new should use this. It differs from ``list_target_types_async`` in exactly one + way: it drops registry-reference parameters, which the un-migrated + configuration UI cannot render. Delete this method, the ``/catalog`` route, and + the ``TargetCatalog*`` aliases together when that UI switches to + ``/targets/types``. + + Returns: + TargetCatalogResponse: The scalar-only legacy projection. + """ + types_response = await self.list_target_types_async() + items = [ + entry.model_copy(update={"parameters": [p for p in entry.parameters if p.is_string_coercible]}) + for entry in types_response.items + ] return TargetCatalogResponse(items=items) async def create_target_async(self, *, request: CreateTargetRequest) -> TargetInstance: @@ -188,11 +212,14 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn # Omit any api_key so the target validates its own endpoint and authenticates itself. params.pop("api_key", None) - target_obj = self._registry.create_instance(request.type, **params) - - self._registry.instances.register(target_obj) - - target_registry_name = target_obj.get_identifier().unique_name + # LEGACY COMPATIBILITY: The current configuration UI omits the name. + # Remove this generated fallback after that UI sends an explicit name. + target_registry_name = request.name or f"compat_{uuid.uuid4().hex}" + target_obj = self._registry.create_named_instance( + name=target_registry_name, + target_type=request.type, + **params, + ) return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) diff --git a/pyrit/converter/add_image_to_video_converter.py b/pyrit/converter/add_image_to_video_converter.py index 376202d283..b59664d396 100644 --- a/pyrit/converter/add_image_to_video_converter.py +++ b/pyrit/converter/add_image_to_video_converter.py @@ -38,7 +38,7 @@ class AddImageVideoConverter(Converter): def __init__( self, *, - video_path: str, + video_path: Path, output_path: str | None = None, img_position: tuple[int, int] = (10, 10), img_resize_size: tuple[int, int] = (500, 500), @@ -47,7 +47,11 @@ def __init__( Initialize the converter with the video path and image properties. Args: - video_path (str): File path of video to add image to. + video_path (Path): File path of the video to add the image to. Declared as a + ``Path`` so it is an input file everywhere it is described (registry + metadata, CLI, REST) instead of a free-form string. It is kept as a + string internally because the memory serializer also accepts an Azure + Blob URL for this value. output_path (str, Optional): File path of output video. Defaults to None. img_position (tuple): Position to place image in video. Defaults to (10, 10). img_resize_size (tuple): Size to resize image to. Defaults to (500, 500). @@ -61,7 +65,7 @@ def __init__( self._output_path = output_path self._img_position = img_position self._img_resize_size = img_resize_size - self._video_path = video_path + self._video_path = str(video_path) def _build_identifier(self) -> ComponentIdentifier: """ diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index ae2da9cd6b..d9304752a5 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -9,14 +9,21 @@ import types from dataclasses import dataclass from enum import Enum +from pathlib import Path 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 -_SUPPORTED_SCALAR_TYPES: tuple[type, ...] = (str, int, float, bool) -_SCALAR_NAME_TO_TYPE: dict[str, type] = {"int": int, "float": float, "bool": bool, "str": str} +_SUPPORTED_SCALAR_TYPES: tuple[type, ...] = (str, int, float, bool, Path) +_SCALAR_NAME_TO_TYPE: dict[str, type] = { + "Path": Path, + "bool": bool, + "float": float, + "int": int, + "str": str, +} class ComponentType(str, Enum): @@ -66,8 +73,8 @@ class Parameter(BaseModel): ``reference``, when set, marks the parameter as a registry reference: its value is supplied *by name* and resolved to a registered instance by the registry - layer (``Parameter`` itself never resolves references). It is also excluded - from serialization. + layer (``Parameter`` itself never resolves references). The live reference is + excluded from serialization; ``reference_type`` exposes its component family. ``coerce_value`` and ``validate`` are the only public behaviors; all coercion branching lives behind them so callers never touch a free function. @@ -130,14 +137,19 @@ def _reconstruct_param_type_from_wire(cls, data: Any) -> Any: """ if not isinstance(data, dict): return data - if data.get("param_type") is not None or "type_name" not in data: + needs_param_type = data.get("param_type") is None and "type_name" in data + needs_reference = data.get("reference") is None and data.get("reference_type") is not None + if not needs_param_type and not needs_reference: return data data = dict(data) - data["param_type"] = _param_type_from_display( - type_name=data.get("type_name"), - choices=data.get("choices"), - is_list=bool(data.get("is_list")), - ) + if needs_param_type: + data["param_type"] = _param_type_from_display( + type_name=data.get("type_name"), + choices=data.get("choices"), + is_list=bool(data.get("is_list")), + ) + if needs_reference: + data["reference"] = RegistryReference(component_type=ComponentType(data["reference_type"])) return data @computed_field @@ -165,6 +177,12 @@ def is_list(self) -> bool: """True when the parameter accepts a list of values (e.g. ``list[str]``).""" return get_origin(self.param_type) is list + @computed_field + @property + 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 + @field_serializer("default") def _serialize_default(self, value: Any) -> str | list[str] | None: """ @@ -191,7 +209,7 @@ def is_string_coercible(self) -> bool: Whether a single string token can be coerced to this parameter's value. True for a non-reference plain scalar (``str`` / ``int`` / ``float`` / - ``bool``), ``Literal[...]``, or ``Enum`` parameter — exactly the forms a + ``bool`` / ``Path``), ``Literal[...]``, or ``Enum`` parameter — exactly the forms a text field or CLI token can supply. References and structured types (lists and arbitrary objects) are False and are surfaced/handled elsewhere. @@ -284,7 +302,7 @@ def validate(self) -> None: # type: ignore[ty:invalid-method-override] raise ValueError( f"Parameter '{self.name}' has unsupported param_type {param_type!r}. " - f"Supported types: str, int, float, bool, Literal[...], Enum, a list of those, " + f"Supported types: str, int, float, bool, Path, Literal[...], Enum, a list of those, " f"or None (or provide a default)." ) @@ -314,7 +332,7 @@ def _is_scalar_param_type(annotation: Any) -> bool: """ Return True when ``annotation`` is a coercible scalar form. - A scalar form is a plain scalar (``str`` / ``int`` / ``float`` / ``bool``) or a + A scalar form is a plain scalar (``str`` / ``int`` / ``float`` / ``bool`` / ``Path``) or a constrained scalar (``Literal[...]`` or an ``Enum`` subclass) that carries its own allowed set. @@ -358,6 +376,8 @@ def _coerce_simple_value(*, param_name: str, annotation: Any, raw_value: Any) -> return _coerce_scalar(param_name=param_name, scalar_type=float, raw_value=raw_value) if annotation is str: return str(raw_value) + if annotation is Path: + return Path(raw_value) return raw_value diff --git a/pyrit/registry/components/converter_registry.py b/pyrit/registry/components/converter_registry.py index 138d4e52ef..60cdd278ac 100644 --- a/pyrit/registry/components/converter_registry.py +++ b/pyrit/registry/components/converter_registry.py @@ -16,15 +16,15 @@ It is a ``Registry``: the registry's own surface (``get_class``, ``get_class_names``, ``get_all_registered_class_metadata``, ``create_instance``) -is the buildable class catalog. Pre-configured instances live under the -``instances`` property (``register``, ``get``, ``get_all_instances``, +is the registered converter-class surface. Pre-configured instances live under the +``instances`` property (``register``, ``get``, ``unregister``, ``get_all_instances``, ``get_names``), a ``DefaultInstanceRegistry``. """ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from pyrit.models.identifiers import ConverterIdentifier from pyrit.models.parameter import ComponentType @@ -76,7 +76,7 @@ class ConverterRegistry(Registry["Converter", ConverterMetadata]): Discovers all concrete ``Converter`` subclasses exported from ``pyrit.converter`` (keyed by their exact class name, e.g. - ``"Base64Converter"``) for the buildable catalog. Pre-configured instances + ``"Base64Converter"``) as registered buildable classes. Pre-configured instances registered via initializers or the backend are held under the ``instances`` property. @@ -94,7 +94,36 @@ def __init__(self, *, lazy_discovery: bool = True) -> None: access. If False, discovery runs immediately. """ super().__init__(lazy_discovery=lazy_discovery) - self.instances: InstanceRegistry[Converter] = DefaultInstanceRegistry(instance_type=self._base_type) + self.instances: InstanceRegistry[Converter] = DefaultInstanceRegistry( + instance_type=self._base_type, + reserved_names={"catalog", "preview", "types"}, + ) + + def create_named_instance( + self, + *, + name: str, + converter_type: str, + registry_metadata: dict[str, Any] | None = None, + **kwargs: object, + ) -> Converter: + """ + Build and store a converter under an explicit registry name. + + Args: + name (str): The unique registry name. + converter_type (str): The registered converter class name. + registry_metadata (dict[str, Any] | None): Per-entry lifecycle metadata + to store with the instance. + **kwargs (object): Constructor arguments. + + Returns: + Converter: The constructed and registered converter. + """ + self.instances.validate_name_available(name) + converter = self.create_instance(converter_type, **kwargs) + self.instances.register(converter, name=name, metadata=registry_metadata) + return converter def _base_type(self) -> type[Converter]: """Return the ``Converter`` base class, imported lazily.""" diff --git a/pyrit/registry/components/target_registry.py b/pyrit/registry/components/target_registry.py index 4cd941f5f6..3aa9d9b3d5 100644 --- a/pyrit/registry/components/target_registry.py +++ b/pyrit/registry/components/target_registry.py @@ -81,7 +81,27 @@ def __init__(self, *, lazy_discovery: bool = True) -> None: access. If False, discovery runs immediately. """ super().__init__(lazy_discovery=lazy_discovery) - self.instances: InstanceRegistry[PromptTarget] = DefaultInstanceRegistry(instance_type=self._base_type) + self.instances: InstanceRegistry[PromptTarget] = DefaultInstanceRegistry( + instance_type=self._base_type, + reserved_names={"catalog", "types"}, + ) + + def create_named_instance(self, *, name: str, target_type: str, **kwargs: object) -> PromptTarget: + """ + Build and store a target under an explicit registry name. + + Args: + name (str): The unique registry name. + target_type (str): The registered target class name. + **kwargs (object): Constructor arguments. + + Returns: + PromptTarget: The constructed and registered target. + """ + self.instances.validate_name_available(name) + target = self.create_instance(target_type, **kwargs) + self.instances.register(target, name=name) + return target def _base_type(self) -> type[PromptTarget]: """Return the ``PromptTarget`` base class, imported lazily.""" diff --git a/pyrit/registry/instance_registry.py b/pyrit/registry/instance_registry.py index 404189e2d1..9cf1f0be0b 100644 --- a/pyrit/registry/instance_registry.py +++ b/pyrit/registry/instance_registry.py @@ -78,6 +78,7 @@ def register( name: str | None = None, tags: dict[str, str] | list[str] | None = None, metadata: dict[str, Any] | None = None, + replace: bool = False, ) -> None: """Register a pre-configured instance, defaulting its name to the identifier's ``unique_name``.""" ... @@ -86,6 +87,14 @@ def get(self, name: str) -> T | None: """Return the instance registered under ``name``, or None.""" ... + def validate_name_available(self, name: str) -> None: + """Raise if ``name`` is reserved or already registered.""" + ... + + def unregister(self, name: str) -> T | None: + """Remove and return the instance registered under ``name``, or None.""" + ... + def get_entry(self, name: str) -> RegistryEntry[T] | None: """Return the full entry (including tags) for ``name``, or None.""" ... @@ -170,7 +179,12 @@ class DefaultInstanceRegistry(Generic[T]): T: The type of instances held (must be ``Identifiable``). """ - def __init__(self, *, instance_type: type[T] | Callable[[], type[T]] | None = None) -> None: + def __init__( + self, + *, + instance_type: type[T] | Callable[[], type[T]] | None = None, + reserved_names: set[str] | frozenset[str] | None = None, + ) -> None: """ Initialize an empty instance container. @@ -183,10 +197,13 @@ def __init__(self, *, instance_type: type[T] | Callable[[], type[T]] | None = No zero-argument callable returning it; the callable form lets owners defer importing the type so a registry's lazy discovery is preserved. It is resolved once, on the first ``register`` call, and cached. + reserved_names (set[str] | frozenset[str] | None): Names that cannot be + registered in this container. """ self._registry_items: dict[str, RegistryEntry[T]] = {} self._metadata_cache: list[ComponentIdentifier] | None = None self._instance_type: type[T] | Callable[[], type[T]] | None = instance_type + self._reserved_names = frozenset(reserved_names or ()) def _resolve_instance_type(self) -> type | None: """ @@ -228,6 +245,7 @@ def register( name: str | None = None, tags: dict[str, str] | list[str] | None = None, metadata: dict[str, Any] | None = None, + replace: bool = False, ) -> None: """ Register a pre-configured instance. @@ -239,10 +257,13 @@ def register( tags (dict[str, str] | list[str] | None): Optional tags for categorization. metadata (dict[str, Any] | None): Optional per-entry metadata. + replace (bool): Whether to replace an existing entry with the same name. Raises: TypeError: If this registry was created with an ``instance_type`` and ``instance`` is not of that type. + ValueError: If the name is reserved or already registered and ``replace`` + is False. """ expected_type = self._resolve_instance_type() if expected_type is not None and not isinstance(instance, expected_type): @@ -253,6 +274,10 @@ def register( if name is None: name = instance.get_identifier().unique_name + if name in self._reserved_names: + raise ValueError(f"Instance name '{name}' is reserved") + if not replace: + self.validate_name_available(name) self._registry_items[name] = RegistryEntry( name=name, @@ -275,6 +300,37 @@ def get(self, name: str) -> T | None: entry = self._registry_items.get(name) return entry.instance if entry is not None else None + def validate_name_available(self, name: str) -> None: + """ + Validate that a registry name can be used. + + Args: + name (str): The proposed registry name. + + Raises: + ValueError: If the name is reserved or already registered. + """ + if name in self._reserved_names: + raise ValueError(f"Instance name '{name}' is reserved") + if name in self._registry_items: + raise ValueError(f"Instance '{name}' already exists") + + def unregister(self, name: str) -> T | None: + """ + Remove a registered instance by name. + + Args: + name (str): The registry name of the instance. + + Returns: + T | None: The removed instance, or None if not found. + """ + entry = self._registry_items.pop(name, None) + if entry is None: + return None + self._metadata_cache = None + return entry.instance + def get_entry(self, name: str) -> RegistryEntry[T] | None: """ Get the full entry (including tags) by name. diff --git a/pyrit/setup/initializers/scorers.py b/pyrit/setup/initializers/scorers.py index 3bcf2e46c9..43b96f1f05 100644 --- a/pyrit/setup/initializers/scorers.py +++ b/pyrit/setup/initializers/scorers.py @@ -703,7 +703,7 @@ def _try_register( try: scorer = factory() - scorer_registry.instances.register(scorer, name=name, tags=list(tags) if tags else None) + scorer_registry.instances.register(scorer, name=name, tags=list(tags) if tags else None, replace=True) logger.info(f"Registered scorer: {name}") except (ValueError, TypeError, KeyError) as e: logger.warning(f"Skipping scorer {name}: {e}") diff --git a/pyrit/setup/initializers/targets.py b/pyrit/setup/initializers/targets.py index 35e491304f..c54c8a56db 100644 --- a/pyrit/setup/initializers/targets.py +++ b/pyrit/setup/initializers/targets.py @@ -703,7 +703,7 @@ def _register_target(self, config: TargetConfig) -> None: target = config.target_class(**kwargs) registry = TargetRegistry.get_registry_singleton() - registry.instances.register(target, name=config.registry_name) + registry.instances.register(target, name=config.registry_name, replace=True) if config.tags: registry.instances.add_tags(name=config.registry_name, tags=list(config.tags)) if config.default_objective_target: @@ -743,12 +743,14 @@ def _configure_adversarial_chat(self) -> None: primary, name="adversarial_chat_primary", tags=[TargetInitializerTags.DEFAULT], + replace=True, ) registry.instances.register( canonical_target, name="adversarial_chat", tags=[TargetInitializerTags.DEFAULT], + replace=True, ) def _auto_group_targets(self) -> None: diff --git a/tests/unit/backend/test_api_routes.py b/tests/unit/backend/test_api_routes.py index 90f3a512d2..bf783ea676 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -31,12 +31,13 @@ ConverterInstance, ConverterInstanceListResponse, ConverterPreviewResponse, - CreateConverterResponse, + ConverterTypeResponse, PreviewStep, ) from pyrit.backend.models.targets import ( TargetCatalogResponse, TargetListResponse, + TargetTypeResponse, ) from pyrit.backend.routes import version as version_routes from pyrit.backend.routes.labels import get_label_options @@ -911,6 +912,20 @@ def test_list_target_catalog(self, client: TestClient) -> None: assert data["items"][0]["target_type"] == "OpenAIChatTarget" assert data["items"][0]["supported_auth_modes"] == ["api_key", "identity"] + def test_list_target_types(self, client: TestClient) -> None: + """Test the primary target type metadata route.""" + with patch("pyrit.backend.routes.targets.get_target_service") as mock_get_service: + mock_service = MagicMock() + mock_service.list_target_types_async = AsyncMock( + return_value=TargetTypeResponse(items=[{"target_type": "TextTarget"}]) + ) + mock_get_service.return_value = mock_service + + response = client.get("/api/targets/types") + + assert response.status_code == status.HTTP_200_OK + assert response.json()["items"][0]["target_type"] == "TextTarget" + def test_create_target_success(self, client: TestClient) -> None: """Test successful target creation.""" with patch("pyrit.backend.routes.targets.get_target_service") as mock_get_service: @@ -947,6 +962,14 @@ def test_create_target_invalid_type(self, client: TestClient) -> None: assert response.status_code == status.HTTP_400_BAD_REQUEST + def test_create_target_rejects_unaddressable_registry_name(self, client: TestClient) -> None: + response = client.post( + "/api/targets", + json={"name": "nested/name", "type": "TextTarget", "params": {}}, + ) + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT + def test_create_target_internal_error(self, client: TestClient) -> None: """Test target creation with internal error returns 500.""" with patch("pyrit.backend.routes.targets.get_target_service") as mock_get_service: @@ -1104,15 +1127,29 @@ def test_list_converter_catalog(self, client: TestClient) -> None: data = response.json() assert data["items"][0]["converter_type"] == "Base64Converter" + def test_list_converter_types(self, client: TestClient) -> None: + """Test the primary converter type metadata route.""" + with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: + mock_service = MagicMock() + mock_service.list_converter_types_async = AsyncMock( + return_value=ConverterTypeResponse(items=[{"converter_type": "Base64Converter"}]) + ) + mock_get_service.return_value = mock_service + + response = client.get("/api/converters/types") + + assert response.status_code == status.HTTP_200_OK + assert response.json()["items"][0]["converter_type"] == "Base64Converter" + def test_create_converter_success(self, client: TestClient) -> None: """Test successful converter instance creation.""" with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: mock_service = MagicMock() mock_service.create_converter_async = AsyncMock( - return_value=CreateConverterResponse( + return_value=ConverterInstance( converter_id="conv-1", - converter_type="Base64Converter", - display_name="My Base64", + identifier=ConverterIdentifier(class_name="Base64Converter"), + is_llm_based=False, ) ) mock_get_service.return_value = mock_service @@ -1125,6 +1162,7 @@ def test_create_converter_success(self, client: TestClient) -> None: assert response.status_code == status.HTTP_201_CREATED data = response.json() assert data["converter_id"] == "conv-1" + assert data["identifier"]["class_name"] == "Base64Converter" def test_create_converter_invalid_type(self, client: TestClient) -> None: """Test converter creation with invalid type.""" @@ -1140,6 +1178,14 @@ def test_create_converter_invalid_type(self, client: TestClient) -> None: assert response.status_code == status.HTTP_400_BAD_REQUEST + def test_create_converter_rejects_unaddressable_registry_name(self, client: TestClient) -> None: + response = client.post( + "/api/converters", + json={"name": "nested/name", "type": "Base64Converter", "params": {}}, + ) + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT + def test_create_converter_internal_error(self, client: TestClient) -> None: """Test converter creation with internal error returns 500.""" with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: @@ -1186,6 +1232,27 @@ def test_get_converter_not_found(self, client: TestClient) -> None: assert response.status_code == status.HTTP_404_NOT_FOUND + def test_delete_converter_success(self, client: TestClient) -> None: + with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: + mock_service = MagicMock() + mock_service.delete_converter_async = AsyncMock(return_value=True) + mock_get_service.return_value = mock_service + + response = client.delete("/api/converters/conv-1") + + assert response.status_code == status.HTTP_204_NO_CONTENT + assert response.content == b"" + + def test_delete_converter_not_found(self, client: TestClient) -> None: + with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: + mock_service = MagicMock() + mock_service.delete_converter_async = AsyncMock(return_value=False) + mock_get_service.return_value = mock_service + + response = client.delete("/api/converters/missing") + + assert response.status_code == status.HTTP_404_NOT_FOUND + def test_preview_conversion_success(self, client: TestClient) -> None: """Test previewing a conversion.""" with patch("pyrit.backend.routes.converters.get_converter_service") as mock_get_service: diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index a144b90967..d824fe03b3 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -66,6 +66,11 @@ def get_vocab(self) -> dict[str, int]: return {word: i for i, word in enumerate(_TOKEN_BIJECTION_VOCAB)} +def _make_data_uri(*, mime_type: str, content: bytes) -> str: + """Build a base64 data URI for constructor-upload tests.""" + return f"data:{mime_type};base64,{base64.b64encode(content).decode('ascii')}" + + @pytest.fixture(autouse=True) def reset_registry(): """Reset the converter registry before each test.""" @@ -161,15 +166,50 @@ async def test_catalog_serializes_parameter_type(self) -> None: caesar_param = next(p for p in caesar_entry.parameters if p.name == "caesar_offset") assert caesar_param.type_name == "int" - async def test_catalog_excludes_non_coercible_params(self) -> None: - """Catalog only surfaces params that can be set from a string (e.g. not the LLM target).""" + async def test_types_include_registry_reference_params(self) -> None: + """Type entries surface target references for registry-backed selection.""" service = ConverterService() - result = await service.list_converter_catalog_async() + result = await service.list_converter_types_async() persuasion_entry = next(item for item in result.items if item.converter_type == "PersuasionConverter") assert persuasion_entry.is_llm_based is True - assert all("Target" not in p.type_name for p in persuasion_entry.parameters) + target_param = next(param for param in persuasion_entry.parameters if param.name == "converter_target") + assert target_param.reference_type == "target" + + async def test_catalog_excludes_registry_reference_params(self) -> None: + """The compatibility catalog preserves the scalar-only form contract.""" + service = ConverterService() + + types_result = await service.list_converter_types_async() + catalog_result = await service.list_converter_catalog_async() + + types_entry = next(item for item in types_result.items if item.converter_type == "PersuasionConverter") + catalog_entry = next(item for item in catalog_result.items if item.converter_type == "PersuasionConverter") + assert any(param.name == "converter_target" for param in types_entry.parameters) + assert all(param.name != "converter_target" for param in catalog_entry.parameters) + assert catalog_entry.parameters == [param for param in types_entry.parameters if param.is_string_coercible] + + async def test_types_include_path_parameters(self) -> None: + """Path parameters derived by the registry remain available through REST.""" + service = ConverterService() + + result = await service.list_converter_types_async() + + transparency_entry = next(item for item in result.items if item.converter_type == "TransparencyAttackConverter") + path_param = next(param for param in transparency_entry.parameters if param.name == "benign_image_path") + assert path_param.required is True + assert path_param.type_name == "Path" + + async def test_every_advertised_file_input_is_a_path_parameter(self) -> None: + """Converter file inputs are declared as ``Path`` so REST treats them as uploads.""" + service = ConverterService() + + result = await service.list_converter_types_async() + + video_entry = next(item for item in result.items if item.converter_type == "AddImageVideoConverter") + video_param = next(param for param in video_entry.parameters if param.name == "video_path") + assert video_param.param_type is Path class TestGetConverter: @@ -237,6 +277,7 @@ async def test_create_converter_raises_for_invalid_type(self) -> None: service = ConverterService() request = CreateConverterRequest( + name="invalid", type="NonExistentConverter", params={}, ) @@ -249,22 +290,23 @@ async def test_create_converter_success(self) -> None: service = ConverterService() request = CreateConverterRequest( + name="my-base64", type="Base64Converter", - display_name="My Base64", params={}, ) result = await service.create_converter_async(request=request) - assert result.converter_id is not None - assert result.converter_type == "Base64Converter" - assert result.display_name == "My Base64" + assert result.converter_id == "my-base64" + assert result.identifier.class_name == "Base64Converter" + assert result.is_llm_based is False async def test_create_converter_registers_in_registry(self) -> None: """Test that create_converter registers object in registry.""" service = ConverterService() request = CreateConverterRequest( + name="base64", type="Base64Converter", params={}, ) @@ -275,87 +317,244 @@ async def test_create_converter_registers_in_registry(self) -> None: converter_obj = service.get_converter_object(converter_id=result.converter_id) assert converter_obj is not None + async def test_create_converter_without_name_preserves_chat_compatibility(self) -> None: + service = ConverterService() -class TestPersistDataUriParams: - """Tests for ConverterService._persist_data_uri_params_async (registry-metadata driven).""" + result = await service.create_converter_async( + request=CreateConverterRequest(type="Base64Converter", params={}), + ) - async def test_persist_data_uri_wraps_path_param(self) -> None: - """A data-URI value for a ``Path``-typed constructor param is persisted and wrapped in Path.""" + assert result.converter_id + assert service.get_converter_object(converter_id=result.converter_id) is not None + + async def test_create_converter_rejects_duplicate_name(self) -> None: service = ConverterService() + original = Base64Converter() + service._registry.instances.register(original, name="shared-name") + request = CreateConverterRequest(name="shared-name", type="CaesarConverter", params={}) - mock_serializer = MagicMock() - mock_serializer.value = "/tmp/persisted.pdf" - mock_serializer.save_data_async = AsyncMock() + with pytest.raises(ValueError, match="already exists"): + await service.create_converter_async(request=request) - params = {"existing_pdf": "data:application/pdf;base64,iVBORw0KGgo="} + assert service.get_converter_object(converter_id="shared-name") is original - with patch( - "pyrit.backend.services.converter_service.data_serializer_factory", - return_value=mock_serializer, - ): - result = await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + @pytest.mark.parametrize("name", ["catalog", "preview", "types"]) + async def test_create_converter_rejects_reserved_route_name(self, name: str) -> None: + service = ConverterService() + request = CreateConverterRequest(name=name, type="Base64Converter", params={}) + + with pytest.raises(ValueError, match="reserved"): + await service.create_converter_async(request=request) - assert result["existing_pdf"] == Path("/tmp/persisted.pdf") - mock_serializer.save_data_async.assert_awaited_once_with(data=base64.b64decode("iVBORw0KGgo=")) - async def test_persist_data_uri_keeps_str_param_as_string(self) -> None: - """A data-URI value for a ``str``-typed constructor param is persisted but left as a string.""" +class TestDeleteConverter: + """Tests for ConverterService.delete_converter_async.""" + + async def test_delete_converter_removes_registered_instance(self) -> None: service = ConverterService() + converter_obj = Base64Converter() + service._registry.instances.register(converter_obj, name="conv-1") - mock_serializer = MagicMock() - mock_serializer.value = "/tmp/words.yaml" - mock_serializer.save_data_async = AsyncMock() + assert await service.delete_converter_async(converter_id="conv-1") is True + assert service.get_converter_object(converter_id="conv-1") is None - params = {"wordswap_path": "data:text/yaml;base64,aGVsbG8="} + async def test_delete_converter_returns_false_when_missing(self) -> None: + service = ConverterService() - with patch( - "pyrit.backend.services.converter_service.data_serializer_factory", - return_value=mock_serializer, + assert await service.delete_converter_async(converter_id="missing") is False + + async def test_delete_converter_removes_only_explicitly_owned_uploads(self, tmp_path: Path) -> None: + service = ConverterService() + data_uri = _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n") + request = CreateConverterRequest(name="pdf", type="PDFConverter", params={"existing_pdf": data_uri}) + + with patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path): + await service.create_converter_async(request=request) + entry = service._registry.instances.get_entry("pdf") + assert entry is not None + owned_path = Path(entry.metadata["owned_artifact_paths"][0]) + assert owned_path.is_file() + + assert await service.delete_converter_async(converter_id="pdf") is True + + assert not owned_path.exists() + + async def test_delete_converter_does_not_infer_ownership_from_instance_paths(self, tmp_path: Path) -> None: + service = ConverterService() + existing_pdf = tmp_path / "caller-owned.pdf" + existing_pdf.write_bytes(b"%PDF-1.4\n") + service._registry.create_named_instance( + name="pdf", + converter_type="PDFConverter", + existing_pdf=existing_pdf, + ) + + assert await service.delete_converter_async(converter_id="pdf") is True + assert existing_pdf.is_file() + + +class TestPersistDataUriParams: + """Tests for ConverterService._persist_data_uri_params_async (registry-metadata driven).""" + + async def test_persist_data_uri_materializes_path_in_managed_local_directory(self, tmp_path: Path) -> None: + """A ``Path`` upload stays local even when CentralMemory storage is not local.""" + service = ConverterService() + params = {"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")} + + with ( + patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path), + patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory, ): - result = await service._persist_data_uri_params_async( + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="PDFConverter", + params=params, + ) + + assert result["existing_pdf"].parent == tmp_path + assert result["existing_pdf"].suffix == ".pdf" + assert result["existing_pdf"].read_bytes() == b"%PDF-1.4\n" + assert owned_paths == [result["existing_pdf"]] + mock_factory.assert_not_called() + + async def test_persist_data_uri_does_not_expand_legacy_string_path_support(self) -> None: + """String path parameters remain outside the managed ``Path`` upload contract.""" + service = ConverterService() + data_uri = _make_data_uri(mime_type="text/yaml", content=b"hello") + params = {"wordswap_path": data_uri} + + with patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory: + result, owned_paths = await service._persist_data_uri_params_async( converter_type="ColloquialWordswapConverter", params=params ) - assert result["wordswap_path"] == "/tmp/words.yaml" - assert not isinstance(result["wordswap_path"], Path) + assert result == params + assert owned_paths == [] + mock_factory.assert_not_called() async def test_persist_data_uri_ignores_param_not_on_converter(self) -> None: """A data-URI value under a name that is not a constructor param is left unchanged.""" service = ConverterService() - with patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory: - result = await service._persist_data_uri_params_async( + result, owned_paths = await service._persist_data_uri_params_async( converter_type="PDFConverter", - params={"not_a_param": "data:application/pdf;base64,iVBORw0KGgo="}, + params={"not_a_param": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")}, ) - assert result == {"not_a_param": "data:application/pdf;base64,iVBORw0KGgo="} + assert result["not_a_param"].startswith("data:application/pdf") + assert owned_paths == [] mock_factory.assert_not_called() async def test_persist_data_uri_noop_for_unregistered_type(self) -> None: """When the converter type has no registry metadata, params pass through untouched.""" service = ConverterService() - params = {"existing_pdf": "data:application/pdf;base64,iVBORw0KGgo="} + params = {"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")} with patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory: - result = await service._persist_data_uri_params_async(converter_type="NonExistentConverter", params=params) + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="NonExistentConverter", params=params + ) assert result == params + assert owned_paths == [] mock_factory.assert_not_called() async def test_persist_data_uri_ignores_non_data_uri_values(self) -> None: - """Values that are not data URIs are left unchanged.""" + """Non-upload values remain unchanged for non-Path parameters.""" service = ConverterService() - params = {"existing_pdf": "/already/a/path.pdf", "font_size": 12} + params = {"font_size": 12} with patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory: - result = await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="PDFConverter", params=params + ) assert result == params + assert owned_paths == [] mock_factory.assert_not_called() + async def test_persist_data_uri_keeps_optional_path_none(self) -> None: + service = ConverterService() + + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="PDFConverter", + params={"existing_pdf": None}, + ) + + assert result == {"existing_pdf": None} + assert owned_paths == [] + + async def test_persist_data_uri_rejects_server_path_for_path_parameter(self) -> None: + service = ConverterService() + + with pytest.raises(ValueError, match="must be uploaded as a data URI"): + await service._persist_data_uri_params_async( + converter_type="PDFConverter", + params={"existing_pdf": "C:\\sensitive\\input.pdf"}, + ) + + @pytest.mark.parametrize( + ("mime_type", "expected_suffix"), + [("text/html", ".html"), ("image/svg+xml", ".svg"), ("application/x-not-real", ".bin")], + ) + async def test_persist_data_uri_stores_any_content_type( + self, mime_type: str, expected_suffix: str, tmp_path: Path + ) -> None: + """Uploads are stored verbatim; restricting content is the media route's job.""" + service = ConverterService() + params = {"existing_pdf": _make_data_uri(mime_type=mime_type, content=b"")} + + with patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path): + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="PDFConverter", params=params + ) + + assert result["existing_pdf"].suffix == expected_suffix + assert result["existing_pdf"].read_bytes() == b"" + assert owned_paths == [result["existing_pdf"]] + + async def test_persist_data_uri_rejects_invalid_base64(self, tmp_path: Path) -> None: + service = ConverterService() + params = {"existing_pdf": "data:application/pdf;base64,not-base64!!"} + + with ( + patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path), + pytest.raises(ValueError, match="invalid base64 data"), + ): + await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + + assert list(tmp_path.iterdir()) == [] + + async def test_persist_data_uri_rejects_non_base64_data_uri(self, tmp_path: Path) -> None: + service = ConverterService() + params = {"existing_pdf": "data:text/plain,hello"} + + with ( + patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path), + pytest.raises(ValueError, match="must be a base64 data URI"), + ): + await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + + assert list(tmp_path.iterdir()) == [] + + async def test_create_converter_cleans_upload_when_construction_fails(self, tmp_path: Path) -> None: + service = ConverterService() + params = { + "existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n"), + "font_color": [256, 0, 0], + } + request = CreateConverterRequest(name="invalid-pdf", type="PDFConverter", params=params) + + with ( + patch("pyrit.backend.services.converter_service._REGISTRY_UPLOAD_DIRECTORY", tmp_path), + pytest.raises(ValueError, match="Invalid font_color"), + ): + await service.create_converter_async(request=request) + + assert service._registry.instances.get("invalid-pdf") is None + assert list(tmp_path.iterdir()) == [] + class TestPreviewConversion: """Tests for ConverterService.preview_conversion method.""" diff --git a/tests/unit/backend/test_mappers.py b/tests/unit/backend/test_mappers.py index ff79eba0e2..f4ee572daa 100644 --- a/tests/unit/backend/test_mappers.py +++ b/tests/unit/backend/test_mappers.py @@ -1918,13 +1918,19 @@ def test_maps_converter_with_identifier(self) -> None: ) converter_obj.get_identifier.return_value = identifier - result = converter_object_to_instance("c-1", converter_obj) + result = converter_object_to_instance( + converter_id="c-1", + converter_obj=converter_obj, + is_llm_based=False, + description="Base64 converter", + ) assert result.converter_id == "c-1" assert result.identifier.class_name == "Base64Converter" assert result.identifier.supported_input_types == ["text"] assert result.identifier.supported_output_types == ["text"] assert result.identifier.params["param1"] == "value1" + assert result.description == "Base64 converter" def test_none_input_output_types_stay_none(self) -> None: """Test that absent supported types stay None on the identifier.""" @@ -1935,7 +1941,12 @@ def test_none_input_output_types_stay_none(self) -> None: ) converter_obj.get_identifier.return_value = identifier - result = converter_object_to_instance("c-1", converter_obj) + result = converter_object_to_instance( + converter_id="c-1", + converter_obj=converter_obj, + is_llm_based=False, + description=None, + ) assert result.identifier.supported_input_types is None assert result.identifier.supported_output_types is None diff --git a/tests/unit/backend/test_media_route.py b/tests/unit/backend/test_media_route.py index 77732cc254..3cc4d931f3 100644 --- a/tests/unit/backend/test_media_route.py +++ b/tests/unit/backend/test_media_route.py @@ -49,6 +49,7 @@ def test_serves_existing_file(self, client: TestClient, _mock_memory: Path) -> N assert response.status_code == 200 assert response.headers["content-type"] == "image/png" + assert response.headers["x-content-type-options"] == "nosniff" assert response.content == b"\x89PNG\r\n\x1a\n" def test_rejects_path_outside_results_directory(self, client: TestClient, _mock_memory: Path) -> None: @@ -162,6 +163,37 @@ def test_rejects_yaml_file(self, client: TestClient, _mock_memory: Path) -> None assert response.status_code == 403 + @pytest.mark.parametrize("extension", [".html", ".svg"]) + def test_serves_active_documents_as_neutralized_downloads( + self, + client: TestClient, + _mock_memory: Path, + extension: str, + ) -> None: + """Active documents are stored and served, but never rendered in this origin.""" + file_path = _mock_memory / "prompt-memory-entries" / f"active{extension}" + file_path.write_text("") + + response = client.get("/api/media", params={"path": str(file_path)}) + + assert response.status_code == 200 + assert response.text == "" + assert response.headers["content-type"] == "application/octet-stream" + assert response.headers["content-disposition"].startswith("attachment;") + assert f"active{extension}" in response.headers["content-disposition"] + assert response.headers["x-content-type-options"] == "nosniff" + + def test_serves_documents_as_attachments(self, client: TestClient, _mock_memory: Path) -> None: + """Allowed documents download instead of rendering in the application origin.""" + file_path = _mock_memory / "prompt-memory-entries" / "document.pdf" + file_path.write_bytes(b"%PDF-1.4\n") + + response = client.get("/api/media", params={"path": str(file_path)}) + + assert response.status_code == 200 + assert response.headers["content-disposition"].startswith("attachment;") + assert response.headers["x-content-type-options"] == "nosniff" + def test_rejects_disallowed_subdirectory(self, client: TestClient, _mock_memory: Path) -> None: """Files in non-allowed subdirectories are rejected.""" other_dir = _mock_memory / "other-stuff" diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index e226d33416..08e8d7ff85 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -253,6 +253,19 @@ async def test_catalog_includes_declarative_auth_facts(self) -> None: assert "api_key" in openai_entry.supported_auth_modes assert "identity" in openai_entry.supported_auth_modes + async def test_types_include_references_while_catalog_preserves_scalar_contract(self) -> None: + service = TargetService() + + types_result = await service.list_target_types_async() + catalog_result = await service.list_target_catalog_async() + + types_entry = next(item for item in types_result.items if item.target_type == "RoundRobinTarget") + catalog_entry = next(item for item in catalog_result.items if item.target_type == "RoundRobinTarget") + targets_parameter = next(param for param in types_entry.parameters if param.name == "targets") + assert targets_parameter.reference_type == "target" + assert all(param.name != "targets" for param in catalog_entry.parameters) + assert catalog_entry.parameters == [param for param in types_entry.parameters if param.is_string_coercible] + async def test_catalog_cold_and_warm_results_are_equal(self) -> None: service = TargetService() @@ -357,6 +370,34 @@ async def test_create_target_success(self, sqlite_instance) -> None: assert result.target_registry_name is not None assert result.identifier.class_name == "TextTarget" + async def test_create_target_uses_explicit_registry_name(self, sqlite_instance) -> None: + service = TargetService() + + result = await service.create_target_async( + request=CreateTargetRequest(name="text-target", type="TextTarget", params={}), + ) + + assert result.target_registry_name == "text-target" + assert service.get_target_object(target_registry_name="text-target") is not None + + async def test_create_target_rejects_duplicate_name(self, sqlite_instance) -> None: + service = TargetService() + service._registry.instances.register(MockPromptTarget(), name="shared-name") + + with pytest.raises(ValueError, match="already exists"): + await service.create_target_async( + request=CreateTargetRequest(name="shared-name", type="TextTarget", params={}), + ) + + @pytest.mark.parametrize("name", ["catalog", "types"]) + async def test_create_target_rejects_reserved_route_name(self, sqlite_instance, name: str) -> None: + service = TargetService() + + with pytest.raises(ValueError, match="reserved"): + await service.create_target_async( + request=CreateTargetRequest(name=name, type="TextTarget", params={}), + ) + async def test_create_target_delegates_construction_to_registry(self, sqlite_instance) -> None: """Every target construction path is owned by the registry.""" service = TargetService() diff --git a/tests/unit/backend/test_target_catalog_concurrency.py b/tests/unit/backend/test_target_types_concurrency.py similarity index 78% rename from tests/unit/backend/test_target_catalog_concurrency.py rename to tests/unit/backend/test_target_types_concurrency.py index d2b7cbe46e..d7a5c5a93c 100644 --- a/tests/unit/backend/test_target_catalog_concurrency.py +++ b/tests/unit/backend/test_target_types_concurrency.py @@ -1,7 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -"""Concurrency regressions for target catalog routes.""" +"""Concurrency regressions for target type routes.""" import asyncio from threading import Event @@ -13,7 +13,7 @@ from pyrit.backend.services.target_service import TargetService -async def test_health_remains_schedulable_during_cold_target_catalog() -> None: +async def test_health_remains_schedulable_during_cold_target_types() -> None: discovery_started = Event() discovery_release = Event() discovery_finished = Event() @@ -31,7 +31,7 @@ def _blocking_metadata_discovery() -> list[object]: patch("pyrit.backend.routes.targets.get_target_service", return_value=service), ): async with AsyncClient(transport=transport, base_url="http://test") as client: - catalog_request = asyncio.create_task(client.get("/api/targets/catalog")) + types_request = asyncio.create_task(client.get("/api/targets/types")) assert await asyncio.to_thread(discovery_started.wait, 5) health_response = await asyncio.wait_for(client.get("/api/health"), timeout=2) @@ -39,6 +39,6 @@ def _blocking_metadata_discovery() -> list[object]: assert health_response.status_code == 200 assert not discovery_finished.is_set() discovery_release.set() - catalog_response = await asyncio.wait_for(catalog_request, timeout=2) + types_response = await asyncio.wait_for(types_request, timeout=2) - assert catalog_response.status_code == 200 + assert types_response.status_code == 200 diff --git a/tests/unit/models/test_parameter.py b/tests/unit/models/test_parameter.py index 624a8d870e..0c4af0e435 100644 --- a/tests/unit/models/test_parameter.py +++ b/tests/unit/models/test_parameter.py @@ -4,6 +4,7 @@ """Unit tests for the unified Parameter model and its coercion methods.""" from enum import Enum +from pathlib import Path from typing import Literal import pytest @@ -84,6 +85,7 @@ def test_scalar_with_default(self) -> None: "required": False, "choices": None, "is_list": False, + "reference_type": None, } def test_excludes_live_only_fields(self) -> None: @@ -93,6 +95,19 @@ def test_excludes_live_only_fields(self) -> None: assert "reference" not in dumped assert "destination" not in dumped + def test_reference_type_serializes_component_family(self) -> None: + parameter = Parameter( + name="target", + description="d", + reference=RegistryReference(component_type=ComponentType.TARGET), + ) + dumped = parameter.model_dump() + restored = Parameter.model_validate(dumped) + + assert dumped["reference_type"] == "target" + assert restored.reference == RegistryReference(component_type=ComponentType.TARGET) + assert restored.reference_type == "target" + def test_required_default_serializes_to_none(self) -> None: p = Parameter(name="mode", description="d", default=REQUIRED_VALUE, param_type=Literal["a", "b"]) dumped = p.model_dump() @@ -134,11 +149,20 @@ def test_optional_scalar_unwraps_to_base_name(self) -> None: assert dumped["type_name"] == "int" + def test_path_round_trip_preserves_coercion(self) -> None: + dumped = Parameter(name="input_path", description="d", param_type=Path).model_dump() + + restored = Parameter.model_validate(dumped) + + assert dumped["type_name"] == "Path" + assert restored.param_type is Path + assert restored.coerce_value("images/input.jpg") == Path("images/input.jpg") + class TestIsScalarParamType: """``_is_scalar_param_type`` recognizes plain and constrained scalars.""" - @pytest.mark.parametrize("annotation", [str, int, float, bool, Literal["a", "b"], _Speed]) + @pytest.mark.parametrize("annotation", [str, int, float, bool, Path, Literal["a", "b"], _Speed]) def test_scalar_forms(self, annotation: object) -> None: assert _is_scalar_param_type(annotation) is True @@ -173,7 +197,7 @@ class TestIsStringCoercible: @pytest.mark.parametrize( "param_type", - [str, int, float, bool, Literal["a", "b"], _Speed, int | None, _Speed | None], + [str, int, float, bool, Path, Literal["a", "b"], _Speed, int | None, _Speed | None], ) def test_coercible_value_types(self, param_type: object) -> None: p = Parameter(name="x", description="d", param_type=param_type) @@ -240,6 +264,10 @@ def test_str_passthrough(self) -> None: p = Parameter(name="s", description="d", param_type=str) assert p.coerce_value("hello") == "hello" + def test_path(self) -> None: + p = Parameter(name="path", description="d", param_type=Path) + assert p.coerce_value("images/input.jpg") == Path("images/input.jpg") + def test_int_invalid_raises(self) -> None: p = Parameter(name="n", description="d", param_type=int) with pytest.raises(ValueError, match="could not be coerced to int"): @@ -363,7 +391,7 @@ class TestValidate: @pytest.mark.parametrize( "param_type", - [None, str, int, float, bool, Literal["a", "b"], _Speed, list[str], list[int], list[Literal["a", "b"]]], + [None, str, int, float, bool, Path, Literal["a", "b"], _Speed, list[str], list[int], list[Literal["a", "b"]]], ) def test_supported_forms_ok(self, param_type: object) -> None: Parameter(name="x", description="d", param_type=param_type).validate() diff --git a/tests/unit/registry/test_converter_registry.py b/tests/unit/registry/test_converter_registry.py index 0dd7c53ee1..bfb02ca4f4 100644 --- a/tests/unit/registry/test_converter_registry.py +++ b/tests/unit/registry/test_converter_registry.py @@ -161,15 +161,39 @@ def test_register_instance_multiple_converters_unique_names(self, registry: Conv assert len(registry.instances) == 2 - def test_register_instance_duplicate_name_overwrites(self, registry: ConverterRegistry): + def test_register_instance_duplicate_name_raises(self, registry: ConverterRegistry): converter1 = MockTextConverter() converter2 = MockImageConverter() registry.instances.register(converter1, name="shared_name") - registry.instances.register(converter2, name="shared_name") - assert len(registry.instances) == 1 - assert registry.instances.get("shared_name") is converter2 + with pytest.raises(ValueError, match="already exists"): + registry.instances.register(converter2, name="shared_name") + + assert registry.instances.get("shared_name") is converter1 + + def test_create_named_instance_builds_and_stores_converter(self, registry: ConverterRegistry): + converter = registry.create_named_instance(name="base64", converter_type="Base64Converter") + + assert isinstance(converter, Base64Converter) + assert registry.instances.get("base64") is converter + + def test_create_named_instance_stores_registry_metadata(self, registry: ConverterRegistry): + converter = registry.create_named_instance( + name="base64", + converter_type="Base64Converter", + registry_metadata={"owned_artifact_paths": ["managed.dat"]}, + ) + + entry = registry.instances.get_entry("base64") + assert entry is not None + assert entry.instance is converter + assert entry.metadata == {"owned_artifact_paths": ["managed.dat"]} + + @pytest.mark.parametrize("name", ["catalog", "preview", "types"]) + def test_create_named_instance_rejects_reserved_name(self, registry: ConverterRegistry, name: str): + with pytest.raises(ValueError, match="reserved"): + registry.create_named_instance(name=name, converter_type="Base64Converter") def test_register_instance_rejects_non_converter(self, registry: ConverterRegistry): class NotAConverter: diff --git a/tests/unit/registry/test_instance_registry.py b/tests/unit/registry/test_instance_registry.py index ec607517e4..2b6118faab 100644 --- a/tests/unit/registry/test_instance_registry.py +++ b/tests/unit/registry/test_instance_registry.py @@ -81,13 +81,26 @@ def test_register_multiple_instances(self, registry: DefaultInstanceRegistry[_Te assert len(registry) == 3 assert registry.get("name2") == "value2" - def test_register_overwrites_existing(self, registry: DefaultInstanceRegistry[_TestItem]): + def test_register_rejects_existing_name(self, registry: DefaultInstanceRegistry[_TestItem]): registry.register(_item("original"), name="name") - registry.register(_item("updated"), name="name") - assert len(registry) == 1 + with pytest.raises(ValueError, match="already exists"): + registry.register(_item("updated"), name="name") + + assert registry.get("name") == "original" + + def test_register_can_explicitly_replace_existing(self, registry: DefaultInstanceRegistry[_TestItem]): + registry.register(_item("original"), name="name") + registry.register(_item("updated"), name="name", replace=True) + assert registry.get("name") == "updated" + def test_register_rejects_reserved_name(self): + registry: DefaultInstanceRegistry[_TestItem] = DefaultInstanceRegistry(reserved_names={"types"}) + + with pytest.raises(ValueError, match="reserved"): + registry.register(_item("value"), name="types") + def test_register_defaults_name_to_identifier_unique_name(self, registry: DefaultInstanceRegistry[_TestItem]): registry.register(_item("value1")) @@ -171,6 +184,28 @@ def test_get_entry_nonexistent_returns_none(self, registry: DefaultInstanceRegis assert registry.get_entry("missing") is None +class TestUnregister: + """Tests for unregistering instances.""" + + def test_unregister_removes_and_returns_instance(self, registry: DefaultInstanceRegistry[_TestItem]) -> None: + item = _item("value1") + registry.register(item, name="name1") + + assert registry.unregister("name1") is item + assert registry.get("name1") is None + + def test_unregister_missing_returns_none(self, registry: DefaultInstanceRegistry[_TestItem]) -> None: + assert registry.unregister("missing") is None + + def test_unregister_invalidates_metadata_cache(self, registry: DefaultInstanceRegistry[_TestItem]) -> None: + registry.register(_item("value1"), name="name1") + assert len(registry.list_metadata()) == 1 + + registry.unregister("name1") + + assert registry.list_metadata() == [] + + class TestGetNamesAndAllInstances: """Tests for get_names and get_all_instances.""" diff --git a/tests/unit/registry/test_scorer_registry.py b/tests/unit/registry/test_scorer_registry.py index 28b97e7a31..ea3588fbe5 100644 --- a/tests/unit/registry/test_scorer_registry.py +++ b/tests/unit/registry/test_scorer_registry.py @@ -172,15 +172,16 @@ def test_register_instance_multiple_scorers_unique_names(self, registry: ScorerR assert len(registry.instances) == 2 - def test_register_instance_duplicate_name_overwrites(self, registry: ScorerRegistry): + def test_register_instance_duplicate_name_raises(self, registry: ScorerRegistry): first = MockTrueFalseScorer() second = MockTrueFalseScorer() registry.instances.register(first, name="same_name") - registry.instances.register(second, name="same_name") - assert len(registry.instances) == 1 - assert registry.instances.get("same_name") is second + with pytest.raises(ValueError, match="already exists"): + registry.instances.register(second, name="same_name") + + assert registry.instances.get("same_name") is first def test_register_instance_rejects_non_scorer(self, registry: ScorerRegistry): class NotAScorer: diff --git a/tests/unit/registry/test_target_registry.py b/tests/unit/registry/test_target_registry.py index 60a78a4e09..de0e534937 100644 --- a/tests/unit/registry/test_target_registry.py +++ b/tests/unit/registry/test_target_registry.py @@ -124,15 +124,31 @@ def test_register_instance_multiple_targets_unique_names(self, registry: TargetR assert len(registry.instances) == 2 - def test_register_instance_duplicate_name_overwrites(self, registry: TargetRegistry): + def test_register_instance_duplicate_name_raises(self, registry: TargetRegistry): first = MockPromptTarget(model_name="first") second = MockPromptTarget(model_name="second") registry.instances.register(first, name="same_name") - registry.instances.register(second, name="same_name") - assert len(registry.instances) == 1 - assert registry.instances.get("same_name") is second + with pytest.raises(ValueError, match="already exists"): + registry.instances.register(second, name="same_name") + + assert registry.instances.get("same_name") is first + + def test_create_named_instance_builds_and_stores_target(self, registry: TargetRegistry): + registry.register_class(MockPromptTarget) + + target = registry.create_named_instance(name="mock", target_type="MockPromptTarget") + + assert isinstance(target, MockPromptTarget) + assert registry.instances.get("mock") is target + + @pytest.mark.parametrize("name", ["catalog", "types"]) + def test_create_named_instance_rejects_reserved_name(self, registry: TargetRegistry, name: str): + registry.register_class(MockPromptTarget) + + with pytest.raises(ValueError, match="reserved"): + registry.create_named_instance(name=name, target_type="MockPromptTarget") def test_register_instance_rejects_non_target(self, registry: TargetRegistry): class NotATarget: diff --git a/tests/unit/setup/test_targets_initializer.py b/tests/unit/setup/test_targets_initializer.py index 1a70801966..08c16f8e1d 100644 --- a/tests/unit/setup/test_targets_initializer.py +++ b/tests/unit/setup/test_targets_initializer.py @@ -696,6 +696,21 @@ async def test_multiple_slots_publish_ordered_round_robin(self, slot_count: int) member_name, _ = self.SLOTS[index] assert registry.instances.get(member_name) is round_robin.inner_targets[index] + async def test_repeated_initialization_replaces_primary_alias(self) -> None: + """Repeated initialization refreshes the canonical and primary targets.""" + from pyrit.prompt_target import RoundRobinTarget + + self._set_slots(0, 1) + initializer = TargetInitializer() + + await initializer.initialize_async() + await initializer.initialize_async() + + registry = TargetRegistry.get_registry_singleton() + round_robin = registry.instances.get("adversarial_chat") + assert isinstance(round_robin, RoundRobinTarget) + assert registry.instances.get("adversarial_chat_primary") is round_robin.inner_targets[0] + async def test_noncontiguous_slots_publish_round_robin_without_inferred_duplicate(self) -> None: """Secondary slots compose directly without producing a generic inferred group.""" from pyrit.prompt_target import RoundRobinTarget