diff --git a/doc/code/registry/0_registry.md b/doc/code/registry/0_registry.md index 84d87ad28e..42f1f415ff 100644 --- a/doc/code/registry/0_registry.md +++ b/doc/code/registry/0_registry.md @@ -54,6 +54,31 @@ show_registry_contents(ScenarioRegistry.get_registry_singleton()) | Instantiation | Caller provides parameters | Pre-configured by initializer | | When to use | Self-contained components with deferred configuration | Components requiring constructor parameters or compositional setup | +## Named Component Construction + +Converter, target, and scorer registries use `InstanceHoldingRegistry` to build +components and store them in their `.instances` registry. Use +`create_named_instance(name=..., type_name=..., params=...)` to build and register +a component in one operation. The instance registry stores objects; it does not +construct them. + +Duplicate names raise `ValueError`. Use `.instances.register(..., replace=True)` +only when replacement is intended. Converter and target registries also reject +reserved route names such as `catalog` and `types`. Use `.instances.unregister(name)` +to remove an instance. + +Constructor annotations define parameter metadata and coercion. Use `Path` for a +local file input. Use `Path | str` when a component also supports a remote URL. +For this union, the registry preserves the supplied type: a `Path` stays a `Path`, +and a string stays a string. It never passes a URL through `Path`. Both union +orders have the same metadata, `type_name: "Path | str"`, including after a JSON +round-trip. Optional forms accept `None` in Python; the display type omits `None`, +as it does for other optional parameters. + +The backend owns file-upload handling and cleanup, not the registry. See the +[registry API migration notes](../../gui/0_gui.md#registry-api-migration-notes) +for the REST contract and temporary compatibility behavior. + ## See Also - [Class Registries](1_class_registry.ipynb) - ScenarioRegistry, InitializerRegistry diff --git a/doc/gui/0_gui.md b/doc/gui/0_gui.md index 9680a3b234..c42e0b281d 100644 --- a/doc/gui/0_gui.md +++ b/doc/gui/0_gui.md @@ -190,6 +190,33 @@ Use **Reload** to discard local edits and fetch the latest source content. Saved --- +## Registry API Migration Notes + +Use `/api/converters/types` and `/api/targets/types` for registry build metadata. +These endpoints return all constructor parameters from the registry, including +lists, unions, and component references. The temporary `/catalog` routes retain +their scalar-only filtering for the current UI. +Create requests should supply an explicit registry `name`. Converter creation +returns the complete `ConverterInstance`; read its type from +`identifier.class_name`, not the old top-level `converter_type` field. Treat +returned IDs as opaque registry names, not UUIDs or identifier hashes. + +Constructor parameters typed as `Path` accept base64 data-URI uploads through REST, +not server filesystem paths. Parameters typed as `Path | str` also accept Azure +Blob URLs. This applies to `AddImageVideoConverter.video_path` and +`ImageOverlayConverter.base_image`. Other local file inputs remain `Path`. +Uploads stay in backend-owned temporary storage until deletion or shutdown, +including with Azure-backed memory. Converter outputs still use configured result +storage. Uploads can contain any file type; the media endpoint renders only +allowlisted image, audio, and video extensions inline. Other files, including PDF, +SVG, HTML, text, and executables, download as `application/octet-stream` attachments. + +**Temporary compatibility, scheduled for removal with the chat migration:** +the `/api/converters/catalog` and `/api/targets/catalog` routes project the same +registry metadata for the current UI. Create requests without a name receive a +generated `compat_...` name. New clients should not depend on these routes or +unnamed creation. + ## Connection Health CoPyRIT monitors the backend connection and shows a status banner: diff --git a/pyrit/backend/main.py b/pyrit/backend/main.py index a639114bbb..39d38dee17 100644 --- a/pyrit/backend/main.py +++ b/pyrit/backend/main.py @@ -39,6 +39,7 @@ version, ) from pyrit.backend.services.configuration_file_service import ConfigurationFileService +from pyrit.backend.services.converter_service import get_converter_service from pyrit.backend.services.environment_file_service import EnvironmentFileService from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH from pyrit.registry import InitializerRegistry @@ -110,7 +111,14 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: # don't emit noise and don't perform filesystem side effects. setup_frontend() - yield + converter_service = await asyncio.to_thread(get_converter_service) + try: + yield + finally: + try: + await converter_service.close_async() + finally: + get_converter_service.cache_clear() app = FastAPI( 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 8d83bb20c2..432bbc2e9c 100644 --- a/pyrit/backend/models/__init__.py +++ b/pyrit/backend/models/__init__.py @@ -52,6 +52,8 @@ ConverterInstanceListResponse, ConverterPreviewRequest, ConverterPreviewResponse, + ConverterTypeEntry, + ConverterTypeResponse, CreateConverterRequest, CreateConverterResponse, PreviewStep, @@ -99,6 +101,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..da93d1e3e1 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -8,6 +8,12 @@ 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 controls +whether the browser renders or downloads them. Storage and download support stay +unrestricted on purpose: any file type is a legitimate attack payload. Only +explicitly allowlisted media types render inline; every other type downloads as +opaque bytes. """ import logging @@ -26,8 +32,9 @@ # Only serve files from known media subdirectories under results_path. _ALLOWED_SUBDIRECTORIES = {"prompt-memory-entries", "seed-prompt-entries"} -# Only serve known media file types (allowlist approach). -_ALLOWED_EXTENSIONS = { +# Only these known-safe media types render inline. Every other extension is +# served as an application/octet-stream attachment. +_INLINE_EXTENSIONS = { # Images ".png", ".jpg", @@ -35,7 +42,6 @@ ".gif", ".bmp", ".webp", - ".svg", ".ico", ".tiff", # Audio @@ -51,12 +57,6 @@ ".mov", ".avi", ".mkv", - # Text / documents - ".txt", - ".md", - ".csv", - ".pdf", - ".html", } @@ -92,10 +92,6 @@ def _validate_media_path(*, path: str, allowed_root: Path) -> Path: if not relative_parts or relative_parts[0] not in _ALLOWED_SUBDIRECTORIES: raise HTTPException(status_code=403, detail="Access denied: path is not in a media subdirectory.") - # Only allow known media file extensions - if real_path.suffix.lower() not in _ALLOWED_EXTENSIONS: - raise HTTPException(status_code=403, detail="Access denied: file type is not allowed.") - return real_path @@ -110,6 +106,11 @@ async def serve_media_async( configured results directory (e.g. ``dbdata/prompt-memory-entries/``) to prevent path traversal attacks and exfiltration of sensitive files. + Upload storage and downloads accept any file type. Extensions in + ``_INLINE_EXTENSIONS`` use their inferred media type and can render inline. + Every other extension is returned as an ``application/octet-stream`` + attachment with ``nosniff`` so the browser downloads rather than renders it. + Args: path: Absolute path to the file. @@ -117,7 +118,7 @@ async def serve_media_async( FileResponse with the file content and inferred MIME type. Raises: - HTTPException 403: If the path is outside the allowed directory or has a blocked extension. + HTTPException 403: If the path is outside the allowed directory. HTTPException 404: If the file does not exist. HTTPException 500: If memory is not initialized. """ @@ -134,8 +135,13 @@ 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() + render_inline = extension in _INLINE_EXTENSIONS + guessed_type, _ = mimetypes.guess_type(validated_path) if render_inline else (None, None) return FileResponse( path=validated_path, - media_type=mime_type or "application/octet-stream", + media_type=guessed_type or "application/octet-stream", + filename=None if render_inline else validated_path.name, + 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 8b3ae286ca..21b3ccdec4 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -12,26 +12,34 @@ - Retrieved from registry (pre-registered at startup or created earlier) """ +import asyncio import base64 +import binascii import mimetypes import uuid +from contextlib import suppress from functools import lru_cache from pathlib import Path +from tempfile import TemporaryDirectory from typing import TYPE_CHECKING, Any +import aiofiles +import aiofiles.os + from pyrit.backend.mappers.converter_mappers import converter_object_to_instance from pyrit.backend.models.converters import ( - ConverterCatalogEntry, ConverterCatalogResponse, ConverterInstance, ConverterInstanceListResponse, ConverterPreviewRequest, ConverterPreviewResponse, + ConverterTypeEntry, + ConverterTypeResponse, CreateConverterRequest, - CreateConverterResponse, PreviewStep, ) from pyrit.backend.services.media_persistence import persist_media_value_async +from pyrit.common.azure_storage import is_azure_blob_uri from pyrit.memory import data_serializer_factory from pyrit.models import PromptDataType from pyrit.registry.components import ConverterRegistry @@ -40,6 +48,10 @@ from pyrit.converter import ConverterResult +_OWNED_ARTIFACT_PATHS_KEY = "owned_artifact_paths" +_DEFAULT_UPLOAD_EXTENSION = ".bin" + + class ConverterService: """ Service for managing converter instances. @@ -51,6 +63,8 @@ class ConverterService: def __init__(self) -> None: """Initialize the converter service.""" self._registry = ConverterRegistry.get_registry_singleton() + self._upload_directory = TemporaryDirectory(prefix="pyrit-registry-uploads-") + self._upload_path = Path(self._upload_directory.name).resolve() def _build_instance_from_object(self, *, converter_id: str, converter_obj: Any) -> ConverterInstance: """ @@ -61,12 +75,30 @@ 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 # ======================================================================== + async def close_async(self) -> None: + """Remove this backend's temporary inputs after requests have stopped.""" + owned_entries = [ + entry + for entry in self._registry.instances.get_all_instances() + if any(path.is_relative_to(self._upload_path) for path in self._get_owned_artifact_paths(entry.metadata)) + ] + await asyncio.to_thread(self._upload_directory.cleanup) + for entry in owned_entries: + self._registry.instances.unregister(entry.name, expected_entry=entry) + async def list_converters_async(self) -> ConverterInstanceListResponse: """ List all converter instances. @@ -80,7 +112,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. @@ -89,20 +121,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=list(metadata.parameters), 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 keeps only string-coercible parameters; + registry references and structured parameters are excluded because the + un-migrated chat UI cannot render them. 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: @@ -126,7 +180,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, expected_entry=entry) is not None + + async def create_converter_async(self, *, request: CreateConverterRequest) -> ConverterInstance: """ Create a new converter instance from API request. @@ -137,26 +206,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, + type_name=request.type, + params=params, + registry_metadata={_OWNED_ARTIFACT_PATHS_KEY: [str(path) for path in owned_paths]}, + ) + except (Exception, asyncio.CancelledError): + 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: @@ -223,15 +302,21 @@ 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. + ``Path | str`` parameters also accept Azure Blob URLs, which pass through + unchanged. Their data-URI uploads use the same local storage. + + Inputs remain local until converter deletion or backend shutdown, even with + Azure-backed memory. Converter outputs still use the configured result storage. The set of constructor parameters (and their types) is sourced from the registry's derived ``Parameter`` metadata rather than re-introspecting the @@ -242,45 +327,106 @@ 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 {} + path_params = ( + { + parameter.name: parameter + for parameter in metadata.parameters + if parameter.is_path or parameter.is_path_or_str + } + 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 name not in path_params: + continue + if value is None: + continue + parameter = path_params[name] + if not isinstance(value, str) or not value.startswith("data:"): + if parameter.is_path_or_str and isinstance(value, str) and is_azure_blob_uri(value): + continue + alternative = " or supplied as an Azure Blob URL" if parameter.is_path_or_str else "" + raise ValueError(f"Path parameter '{name}' must be uploaded as a data URI{alternative}") + + content, extension = self._decode_data_uri(parameter_name=name, data_uri=value) + file_path = self._upload_path / f"{uuid.uuid4().hex}{extension}" + async with aiofiles.open(file_path, "xb") as file: + owned_paths.append(file_path) + await file.write(content) result[name] = file_path + except (Exception, asyncio.CancelledError): + await self._remove_owned_artifacts_async(paths=owned_paths) + raise + + return result, owned_paths - return result + @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 + 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] + + async def _remove_owned_artifacts_async(self, *, paths: list[Path]) -> None: + """Remove explicitly owned files, limited to the managed upload directory.""" + for path in paths: + resolved_path = await asyncio.to_thread(path.resolve) + try: + resolved_path.relative_to(self._upload_path) + 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 114001cb10..f088a176dc 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 @@ -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 @@ -147,7 +149,7 @@ def _get_catalog_auth_modes(auth_modes: tuple[str, ...]) -> list[Literal["api_ke raise ValueError(f"Unsupported target authentication mode: {auth_mode!r}") return catalog_auth_modes - 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. @@ -158,18 +160,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=list(metadata.parameters), supported_auth_modes=self._get_catalog_auth_modes(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 keeps only string-coercible parameters; registry + references and structured parameters are excluded because the un-migrated + configuration UI cannot render them. 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: @@ -210,11 +234,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, + type_name=request.type, + params=params, + ) return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) diff --git a/pyrit/converter/add_image_text_converter.py b/pyrit/converter/add_image_text_converter.py index 913ad63b43..117b578b36 100644 --- a/pyrit/converter/add_image_text_converter.py +++ b/pyrit/converter/add_image_text_converter.py @@ -5,6 +5,7 @@ import logging import math from io import BytesIO +from pathlib import Path from typing import cast from PIL import Image, ImageFont @@ -39,8 +40,8 @@ class AddImageTextConverter(_BaseImageTextConverter): def __init__( self, *, - img_to_add: str, - font_name: str | None = None, + img_to_add: Path, + font_name: Path | None = None, color: tuple[int, int, int] = (0, 0, 0), font_size: int | tuple[int, int] = 15, bounding_box: tuple[int, int, int, int] | None = None, @@ -51,8 +52,8 @@ def __init__( Initialize the converter with the image file path and text properties. Args: - img_to_add (str): File path of image to add text to. - font_name (str | None): Path of font to use. Must be a TrueType font (.ttf). + img_to_add (Path): File path of image to add text to. + font_name (Path | None): Path of font to use. Must be a TrueType font (.ttf). Defaults to None which uses Pillow's built-in default font. color (tuple[int, int, int]): Color to print text in, using RGB values. Defaults to (0, 0, 0). font_size (int | tuple[int, int]): Font size as a fixed int, or a (min, max) tuple for automatic @@ -71,7 +72,7 @@ def __init__( """ if not img_to_add: raise ValueError("Please provide valid image path") - if font_name is not None and not font_name.endswith(".ttf"): + if font_name is not None and Path(font_name).suffix.lower() != ".ttf": raise ValueError("The specified font must be a TrueType font with a .ttf extension") self._extract_font_size(font_size) if bounding_box is not None: @@ -80,8 +81,8 @@ def __init__( raise ValueError("bounding_box must have x2 > x1 and y2 > y1") if not math.isfinite(rotation): raise ValueError(f"rotation must be finite, got {rotation}") - self._img_to_add = img_to_add - self._font_name = font_name + self._img_to_add = str(img_to_add) + self._font_name = str(font_name) if font_name is not None else None self._font_size = self._font_size_max self._font_load_failed = font_name is None self._font = self._load_font() diff --git a/pyrit/converter/add_image_to_video_converter.py b/pyrit/converter/add_image_to_video_converter.py index bf9eb9b2dc..9da566e5c6 100644 --- a/pyrit/converter/add_image_to_video_converter.py +++ b/pyrit/converter/add_image_to_video_converter.py @@ -40,7 +40,7 @@ class AddImageVideoConverter(Converter): def __init__( self, *, - video_path: str, + video_path: Path | str, img_position: tuple[int, int] = (10, 10), img_resize_size: tuple[int, int] = (500, 500), ) -> None: @@ -48,7 +48,7 @@ def __init__( Initialize the converter with the video path and image properties. Args: - video_path (str): File path or Azure Blob URL of video to add image to. + video_path (Path | str): Local file path or Azure Blob URL of video to add image to. 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). @@ -60,7 +60,7 @@ def __init__( 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/converter/add_text_image_converter.py b/pyrit/converter/add_text_image_converter.py index c65fc7ef2b..0ee8fa1c01 100644 --- a/pyrit/converter/add_text_image_converter.py +++ b/pyrit/converter/add_text_image_converter.py @@ -5,6 +5,7 @@ import hashlib import logging from io import BytesIO +from pathlib import Path from typing import cast from PIL import Image, ImageFont @@ -34,7 +35,7 @@ def __init__( self, *, text_to_add: str, - font_name: str | None = None, + font_name: Path | None = None, color: tuple[int, int, int] = (0, 0, 0), font_size: int = 15, x_pos: int = 10, @@ -45,7 +46,7 @@ def __init__( Args: text_to_add (str): Text to add to an image. - font_name (str | None): Path of font to use. Must be a TrueType font (.ttf). + font_name (Path | None): Path of font to use. Must be a TrueType font (.ttf). Defaults to None which uses Pillow's built-in default font. color (tuple): Color to print text in, using RGB values. Defaults to (0, 0, 0). font_size (int): Size of font to use. Defaults to 15. @@ -57,10 +58,10 @@ def __init__( """ if text_to_add.strip() == "": raise ValueError("Please provide valid text_to_add value") - if font_name is not None and not font_name.endswith(".ttf"): + if font_name is not None and Path(font_name).suffix.lower() != ".ttf": raise ValueError("The specified font must be a TrueType font with a .ttf extension") self._text_to_add = text_to_add - self._font_name = font_name + self._font_name = str(font_name) if font_name is not None else None self._font_size = font_size self._font = self._load_font() self._color = color diff --git a/pyrit/converter/colloquial_wordswap_converter.py b/pyrit/converter/colloquial_wordswap_converter.py index 8fc6c3eca2..5b4add48af 100644 --- a/pyrit/converter/colloquial_wordswap_converter.py +++ b/pyrit/converter/colloquial_wordswap_converter.py @@ -1,8 +1,8 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -import pathlib import re +from pathlib import Path import yaml @@ -27,7 +27,7 @@ def __init__( *, deterministic: bool = False, custom_substitutions: dict[str, list[str]] | None = None, - wordswap_path: str | None = None, + wordswap_path: Path | None = None, ) -> None: """ Initialize the converter with optional deterministic mode and substitutions source. @@ -37,7 +37,7 @@ def __init__( If False, randomly choose a substitution for each wordswap. Defaults to False. custom_substitutions (dict[str, list[str]] | None): A dictionary of custom substitutions to override the defaults. Defaults to None. - wordswap_path (str | None): Path to a YAML file containing word substitutions. + wordswap_path (Path | None): Path to a YAML file containing word substitutions. Can be a filename within the built-in colloquial_wordswaps directory (e.g., "filipino.yaml") or an absolute path to a custom YAML file. Defaults to None (uses singaporean.yaml). @@ -49,7 +49,7 @@ def __init__( if custom_substitutions is not None and wordswap_path is not None: raise ValueError("Provide either custom_substitutions or wordswap_path, not both.") - self._wordswap_path = wordswap_path + self._wordswap_path = str(wordswap_path) if wordswap_path is not None else None if custom_substitutions is not None and len(custom_substitutions) > 0: self._colloquial_substitutions = custom_substitutions @@ -57,7 +57,7 @@ def __init__( wordswap_directory = CONVERTER_SEED_PROMPT_PATH / "colloquial_wordswaps" if wordswap_path is not None: - file_path = pathlib.Path(wordswap_path) + file_path = Path(wordswap_path) if not file_path.is_absolute(): file_path = wordswap_directory / wordswap_path else: diff --git a/pyrit/converter/image_overlay_converter.py b/pyrit/converter/image_overlay_converter.py index 380a1fbe05..01ffa17d4c 100644 --- a/pyrit/converter/image_overlay_converter.py +++ b/pyrit/converter/image_overlay_converter.py @@ -4,6 +4,7 @@ import base64 import logging from io import BytesIO +from pathlib import Path from PIL import Image @@ -31,7 +32,7 @@ class ImageOverlayConverter(Converter): def __init__( self, *, - base_image: str, + base_image: Path | str, position: tuple[int, int] = (0, 0), overlay_size: tuple[int, int] | None = None, opacity: float = 1.0, @@ -40,7 +41,7 @@ def __init__( Initialize the converter with base image and placement parameters. Args: - base_image (str): File path of the base image onto which overlays will be placed. + base_image (Path | str): Local file path or Azure Blob URL of the base image. position (tuple[int, int]): (x, y) pixel coordinates on the base image where the top-left corner of the overlay will be placed. Defaults to (0, 0). overlay_size (tuple[int, int] | None): Optional (width, height) to resize the @@ -59,7 +60,7 @@ def __init__( if overlay_size is not None and (len(overlay_size) != 2 or overlay_size[0] <= 0 or overlay_size[1] <= 0): raise ValueError("overlay_size must be a tuple of two positive integers (width, height)") - self._base_image = base_image + self._base_image = str(base_image) self._position = position self._overlay_size = overlay_size self._opacity = opacity diff --git a/pyrit/converter/image_prompt_style_converter.py b/pyrit/converter/image_prompt_style_converter.py index f7ee38d731..e03f8e0519 100644 --- a/pyrit/converter/image_prompt_style_converter.py +++ b/pyrit/converter/image_prompt_style_converter.py @@ -40,7 +40,7 @@ def __init__( *, converter_target: PromptTarget = REQUIRED_VALUE, # type: ignore[ty:invalid-parameter-default] filter_name: str | None = None, - filter_path: str | Path | None = None, + filter_path: Path | None = None, variation: str | None = None, ) -> None: """ diff --git a/pyrit/datasets/seed_datasets/remote/comic_jailbreak_dataset.py b/pyrit/datasets/seed_datasets/remote/comic_jailbreak_dataset.py index 865122bc79..72ec1afade 100644 --- a/pyrit/datasets/seed_datasets/remote/comic_jailbreak_dataset.py +++ b/pyrit/datasets/seed_datasets/remote/comic_jailbreak_dataset.py @@ -4,6 +4,7 @@ import logging import uuid from dataclasses import dataclass +from pathlib import Path from typing import TYPE_CHECKING, ClassVar, Literal from typing_extensions import override @@ -362,7 +363,7 @@ async def _render_comic_async( from pyrit.converter import AddImageTextConverter converter = AddImageTextConverter( - img_to_add=template_path, + img_to_add=Path(template_path), bounding_box=bounding_box, rotation=float(rotation), center_text=True, diff --git a/pyrit/models/parameter.py b/pyrit/models/parameter.py index ae2da9cd6b..d3604b4f73 100644 --- a/pyrit/models/parameter.py +++ b/pyrit/models/parameter.py @@ -9,14 +9,22 @@ 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 | types.UnionType] = { + "Path": Path, + "Path | str": Path | str, + "bool": bool, + "float": float, + "int": int, + "str": str, +} class ComponentType(str, Enum): @@ -66,8 +74,10 @@ 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, + while ``type_name`` and ``is_list`` expose whether clients supply one name or a + list of names. ``coerce_value`` and ``validate`` are the only public behaviors; all coercion branching lives behind them so callers never touch a free function. @@ -130,21 +140,37 @@ 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"]), + annotation=data.get("param_type"), + ) return data + @property + def _display_type(self) -> Any: + """Wire type, where registry references are supplied as names.""" + if self.reference is None: + return self.param_type + annotation = _unwrap_optional(self.reference.annotation) + return list[str] if get_origin(annotation) is list else str + @computed_field @property def type_name(self) -> str: """Display name of the parameter's type (e.g. ``'int'``, ``'str'``, ``'list[str]'``, ``'any'``).""" - return _render_type_name(self.param_type) + return _render_type_name(self._display_type) @computed_field @property @@ -156,14 +182,30 @@ def required(self) -> bool: @property def choices(self) -> list[str] | None: """Allowed values for a constrained scalar (``Literal`` / ``Enum``), or None when unconstrained.""" - members = display_choices(self.param_type) + members = display_choices(self._display_type) return [str(member) for member in members] if members is not None else None @computed_field @property 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 + return get_origin(self._display_type) is list + + @property + def is_path(self) -> bool: + """Whether this is a local filesystem path parameter.""" + return self.reference is None and _unwrap_optional(self.param_type) is Path + + @property + def is_path_or_str(self) -> bool: + """Whether this parameter accepts paths or strings without URL normalization.""" + return self.reference is None and _is_path_or_str(self.param_type) + + @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 +233,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`` / ``Path | str``), ``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 +326,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, Path | str, Literal[...], Enum, a list of those, " f"or None (or provide a default)." ) @@ -310,18 +352,26 @@ def _is_enum_type(annotation: Any) -> bool: return isinstance(annotation, type) and issubclass(annotation, Enum) +def _is_path_or_str(annotation: Any) -> bool: + """Return whether the annotation is ``Path | str``, optionally including None.""" + return get_origin(annotation) in (Union, types.UnionType) and set(get_args(annotation)) in ( + {Path, str}, + {Path, str, type(None)}, + ) + + 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. Returns: bool: True when the annotation is a single coercible scalar form. """ - if annotation in _SUPPORTED_SCALAR_TYPES: + if annotation in _SUPPORTED_SCALAR_TYPES or _is_path_or_str(annotation): return True if get_origin(annotation) is Literal: return True @@ -346,6 +396,10 @@ def _coerce_simple_value(*, param_name: str, annotation: Any, raw_value: Any) -> cannot be coerced to the annotated scalar type. """ annotation = _unwrap_optional(annotation) + if _is_path_or_str(annotation): + if isinstance(raw_value, (Path, str)): + return raw_value + raise ValueError(f"Parameter '{param_name}' expects a Path or str, got {type(raw_value).__name__}.") if get_origin(annotation) is Literal: return _coerce_literal(param_name=param_name, annotation=annotation, raw_value=raw_value) if _is_enum_type(annotation): @@ -358,6 +412,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 @@ -504,7 +560,7 @@ def _param_type_from_display(*, type_name: str | None, choices: list[str] | None if not type_name or type_name == "any": return None base_name = type_name.removeprefix("list[").rstrip("]") if is_list else type_name - base_type: type = _SCALAR_NAME_TO_TYPE.get(base_name, str) + base_type = _SCALAR_NAME_TO_TYPE.get(base_name, str) if choices: coerced = tuple(_coerce_simple_value(param_name="", annotation=base_type, raw_value=c) for c in choices) element_type: Any = Literal[coerced] # ty: ignore[invalid-type-form] @@ -533,6 +589,8 @@ def _render_type_name(param_type: Any) -> str: if param_type is None: return "any" param_type = _unwrap_optional(param_type) + if _is_path_or_str(param_type): + return "Path | str" if get_origin(param_type) is Literal: args = get_args(param_type) literal_type_name: str = type(args[0]).__name__ if args else "str" @@ -548,7 +606,7 @@ def _render_type_name(param_type: Any) -> str: member = next(iter(element_type), None) return f"list[{type(member.value).__name__ if member is not None else 'str'}]" if _is_scalar_param_type(element_type): - return f"list[{element_type.__name__}]" + return f"list[{_render_type_name(element_type)}]" # Detect parameterized generics (list[str], dict[str, int], ...) reliably across Python # versions: get_origin returns the unparameterized type for GenericAlias, None otherwise. if get_origin(param_type) is not None: diff --git a/pyrit/registry/__init__.py b/pyrit/registry/__init__.py index 9276f3839f..ed782684e7 100644 --- a/pyrit/registry/__init__.py +++ b/pyrit/registry/__init__.py @@ -30,7 +30,7 @@ RegistryEntry, SupportsInstances, ) - from pyrit.registry.registry import ParamBagRegistry, Registry + from pyrit.registry.registry import InstanceHoldingRegistry, ParamBagRegistry, Registry from pyrit.registry.registry_metadata import RegistryMetadata from pyrit.registry.tag_query import TagQuery @@ -41,6 +41,7 @@ "ConverterMetadata": "pyrit.registry.components", "DefaultInstanceRegistry": "pyrit.registry.instance_registry", "InstanceRegistry": "pyrit.registry.instance_registry", + "InstanceHoldingRegistry": "pyrit.registry.registry", "ParamBagRegistry": "pyrit.registry.registry", "Registry": "pyrit.registry.registry", "RegistryMetadata": "pyrit.registry.registry_metadata", diff --git a/pyrit/registry/components/converter_registry.py b/pyrit/registry/components/converter_registry.py index 138d4e52ef..4510aef96b 100644 --- a/pyrit/registry/components/converter_registry.py +++ b/pyrit/registry/components/converter_registry.py @@ -16,8 +16,8 @@ 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``. """ @@ -28,8 +28,7 @@ from pyrit.models.identifiers import ConverterIdentifier from pyrit.models.parameter import ComponentType -from pyrit.registry.instance_registry import DefaultInstanceRegistry, InstanceRegistry -from pyrit.registry.registry import Registry +from pyrit.registry.registry import InstanceHoldingRegistry from pyrit.registry.registry_metadata import RegistryMetadata if TYPE_CHECKING: @@ -70,13 +69,13 @@ def is_llm_based(self) -> bool: return any(p.is_reference_to(ComponentType.TARGET) for p in self.parameters) -class ConverterRegistry(Registry["Converter", ConverterMetadata]): +class ConverterRegistry(InstanceHoldingRegistry["Converter", ConverterMetadata]): """ Registry that discovers, builds, and holds ``Converter`` instances. 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. @@ -93,8 +92,10 @@ def __init__(self, *, lazy_discovery: bool = True) -> None: lazy_discovery (bool): If True, class discovery is deferred until first access. If False, discovery runs immediately. """ - super().__init__(lazy_discovery=lazy_discovery) - self.instances: InstanceRegistry[Converter] = DefaultInstanceRegistry(instance_type=self._base_type) + super().__init__( + lazy_discovery=lazy_discovery, + reserved_instance_names={"catalog", "preview", "types"}, + ) def _base_type(self) -> type[Converter]: """Return the ``Converter`` base class, imported lazily.""" diff --git a/pyrit/registry/components/scorer_registry.py b/pyrit/registry/components/scorer_registry.py index 558311084f..a2fadb1c7f 100644 --- a/pyrit/registry/components/scorer_registry.py +++ b/pyrit/registry/components/scorer_registry.py @@ -29,8 +29,7 @@ from pyrit.models.identifiers import ScorerIdentifier from pyrit.models.parameter import ComponentType -from pyrit.registry.instance_registry import DefaultInstanceRegistry, InstanceRegistry -from pyrit.registry.registry import Registry +from pyrit.registry.registry import InstanceHoldingRegistry from pyrit.registry.registry_metadata import RegistryMetadata if TYPE_CHECKING: @@ -59,7 +58,7 @@ def is_llm_based(self) -> bool: return any(p.is_reference_to(ComponentType.TARGET) for p in self.parameters) -class ScorerRegistry(Registry["Scorer", ScorerMetadata]): +class ScorerRegistry(InstanceHoldingRegistry["Scorer", ScorerMetadata]): """ Registry that discovers, builds, and holds ``Scorer`` instances. @@ -83,7 +82,6 @@ def __init__(self, *, lazy_discovery: bool = True) -> None: access. If False, discovery runs immediately. """ super().__init__(lazy_discovery=lazy_discovery) - self.instances: InstanceRegistry[Scorer] = DefaultInstanceRegistry(instance_type=self._base_type) def _base_type(self) -> type[Scorer]: """Return the ``Scorer`` base class, imported lazily.""" diff --git a/pyrit/registry/components/target_registry.py b/pyrit/registry/components/target_registry.py index 4cd941f5f6..61509c9238 100644 --- a/pyrit/registry/components/target_registry.py +++ b/pyrit/registry/components/target_registry.py @@ -27,8 +27,7 @@ from typing import TYPE_CHECKING from pyrit.models.identifiers import TargetIdentifier -from pyrit.registry.instance_registry import DefaultInstanceRegistry, InstanceRegistry -from pyrit.registry.registry import Registry +from pyrit.registry.registry import InstanceHoldingRegistry from pyrit.registry.registry_metadata import RegistryMetadata if TYPE_CHECKING: @@ -57,7 +56,7 @@ def supported_auth_modes(self) -> tuple[str, ...]: return auth_modes -class TargetRegistry(Registry["PromptTarget", TargetMetadata]): +class TargetRegistry(InstanceHoldingRegistry["PromptTarget", TargetMetadata]): """ Registry that discovers, builds, and holds ``PromptTarget`` instances. @@ -80,8 +79,10 @@ def __init__(self, *, lazy_discovery: bool = True) -> None: lazy_discovery (bool): If True, class discovery is deferred until first access. If False, discovery runs immediately. """ - super().__init__(lazy_discovery=lazy_discovery) - self.instances: InstanceRegistry[PromptTarget] = DefaultInstanceRegistry(instance_type=self._base_type) + super().__init__( + lazy_discovery=lazy_discovery, + reserved_instance_names={"catalog", "types"}, + ) 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..6e9e5440d5 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, *, expected_entry: RegistryEntry[T] | None = None) -> T | None: + """Remove and return ``name`` when it still refers to ``expected_entry``, if supplied.""" + ... + 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,41 @@ 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, *, expected_entry: RegistryEntry[T] | None = None) -> T | None: + """ + Remove a registered instance by name. + + Args: + name (str): The registry name of the instance. + expected_entry (RegistryEntry[T] | None): When supplied, remove the + name only if it still refers to this exact registry entry. + + Returns: + T | None: The removed instance, or None if the name is missing or now + refers to a different entry. + """ + entry = self._registry_items.get(name) + if entry is None or (expected_entry is not None and entry is not expected_entry): + return None + del self._registry_items[name] + 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/registry/registry.py b/pyrit/registry/registry.py index 43674cafe7..9a259dce3b 100644 --- a/pyrit/registry/registry.py +++ b/pyrit/registry/registry.py @@ -32,6 +32,7 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar +from pyrit.registry.instance_registry import DefaultInstanceRegistry, InstanceRegistry from pyrit.registry.registry_metadata import RegistryMetadata from pyrit.registry.resolution import ( derive_parameters, @@ -43,12 +44,14 @@ from types import ModuleType from typing import Self + from pyrit.models import Identifiable from pyrit.models.identifiers.component_identifier import ComponentIdentifier from pyrit.models.parameter import ComponentType, Parameter logger = logging.getLogger(__name__) T = TypeVar("T") +InstanceT = TypeVar("InstanceT", bound="Identifiable") MetadataT = TypeVar("MetadataT", bound=RegistryMetadata) ConfigurableT = TypeVar("ConfigurableT", bound="SupportsParamBag") @@ -736,6 +739,68 @@ def __iter__(self) -> Iterator[str]: return iter(self.get_class_names()) +class InstanceHoldingRegistry(Registry[InstanceT, MetadataT]): + """ + Registry that builds classes and stores named, configured instances. + + Extends the class-catalog and construction behavior of ``Registry`` with a + typed ``instances`` container. This keeps construction on the owning component + registry while ``InstanceRegistry`` remains responsible only for storing and + retrieving already-built objects. + + Type Parameters: + InstanceT: The identifiable component type that this registry builds and stores. + MetadataT: The metadata dataclass for buildable classes. + """ + + def __init__( + self, + *, + lazy_discovery: bool = True, + reserved_instance_names: set[str] | frozenset[str] | None = None, + ) -> None: + """ + Initialize the class catalog and typed instance container. + + Args: + lazy_discovery (bool): If True, class discovery is deferred until first + access. If False, discovery runs immediately. + reserved_instance_names (set[str] | frozenset[str] | None): Names that + cannot be used for stored instances. + """ + super().__init__(lazy_discovery=lazy_discovery) + self.instances: InstanceRegistry[InstanceT] = DefaultInstanceRegistry( + instance_type=self._base_type, + reserved_names=reserved_instance_names, + ) + + def create_named_instance( + self, + *, + name: str, + type_name: str, + params: Mapping[str, object] | None = None, + registry_metadata: dict[str, Any] | None = None, + ) -> InstanceT: + """ + Build and store a configured instance under an explicit name. + + Args: + name (str): The unique instance name. + type_name (str): The registered class name to build. + params (Mapping[str, object] | None): Constructor arguments. + registry_metadata (dict[str, Any] | None): Per-entry metadata to store + with the instance. + + Returns: + InstanceT: The constructed and registered instance. + """ + self.instances.validate_name_available(name) + instance = self.create_instance(type_name, **dict(params) if params is not None else {}) + self.instances.register(instance, name=name, metadata=registry_metadata) + return instance + + class ParamBagRegistry(Registry[ConfigurableT, MetadataT]): """ Registry whose components carry a parameter bag populated post-construction. diff --git a/pyrit/setup/initializers/converters.py b/pyrit/setup/initializers/converters.py index 93034295cd..0c2ccd2988 100644 --- a/pyrit/setup/initializers/converters.py +++ b/pyrit/setup/initializers/converters.py @@ -131,7 +131,7 @@ class ConverterInitializer(PyRITInitializer): registry_name="add_image_text", converter_type="AddImageTextConverter", constructor_args={ - "img_to_add": str(DATASETS_PATH / "seed_datasets" / "local" / "examples" / "blank_canvas.png") + "img_to_add": DATASETS_PATH / "seed_datasets" / "local" / "examples" / "blank_canvas.png" }, ), ConverterConfig( @@ -167,7 +167,7 @@ async def initialize_async(self) -> None: converter_registry=converter_registry, config=config, ) - converter_registry.instances.register(converter, name=config.registry_name) + converter_registry.instances.register(converter, name=config.registry_name, replace=True) logger.info("Registered converter: %s", config.registry_name) except (FileNotFoundError, KeyError, TypeError, ValueError) as ex: logger.warning("Skipping converter '%s': %s", config.registry_name, ex) 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/pyrit/setup/initializers/techniques/airt.py b/pyrit/setup/initializers/techniques/airt.py index 0ca857c6fd..c372645004 100644 --- a/pyrit/setup/initializers/techniques/airt.py +++ b/pyrit/setup/initializers/techniques/airt.py @@ -22,8 +22,6 @@ from pyrit.prompt_normalizer import ConverterConfiguration from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory -_BLANK_IMAGE_PATH = str(DATASETS_PATH / "seed_datasets" / "local" / "examples" / "blank_canvas.png") - def get_technique_factories() -> list[AttackTechniqueFactory]: """ @@ -57,7 +55,11 @@ def get_technique_factories() -> list[AttackTechniqueFactory]: attack_kwargs={ "attack_converter_config": AttackConverterConfig( request_converters=ConverterConfiguration.from_converters( - converters=[AddImageTextConverter(img_to_add=_BLANK_IMAGE_PATH)] + converters=[ + AddImageTextConverter( + img_to_add=DATASETS_PATH / "seed_datasets" / "local" / "examples" / "blank_canvas.png" + ) + ] ) ), }, diff --git a/tests/unit/backend/test_api_routes.py b/tests/unit/backend/test_api_routes.py index c808a861cf..fe420eee37 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -34,12 +34,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.scores import _get_user_identifier @@ -1090,6 +1091,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: @@ -1126,6 +1141,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: @@ -1283,15 +1306,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 @@ -1304,6 +1341,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.""" @@ -1319,6 +1357,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: @@ -1365,6 +1411,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 50c9e9c362..75dc755aea 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -5,7 +5,9 @@ Tests for backend converter service. """ +import asyncio import base64 +from collections.abc import AsyncGenerator from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch @@ -27,6 +29,7 @@ SuffixAppendConverter, ) from pyrit.converter.converter import get_converter_modalities +from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ComponentIdentifier from pyrit.registry.components import ConverterRegistry @@ -66,6 +69,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.""" @@ -74,6 +82,15 @@ def reset_registry(): ConverterRegistry.reset_registry_singleton() +@pytest.fixture +async def upload_service() -> AsyncGenerator[ConverterService, None]: + service = ConverterService() + try: + yield service + finally: + await service.close_async() + + class TestListConverters: """Tests for ConverterService.list_converters method.""" @@ -162,25 +179,123 @@ async def test_catalog_serializes_parameter_type(self) -> None: assert caesar_param.type_name == "int" async def test_catalog_exposes_video_input_without_output_path(self) -> None: - """The video converter accepts an uploaded input but no caller-controlled destination.""" + """The video converter accepts a local path or URL but no caller-controlled destination.""" service = ConverterService() result = await service.list_converter_catalog_async() video_entry = next(item for item in result.items if item.converter_type == "AddImageVideoConverter") video_path_param = next(parameter for parameter in video_entry.parameters if parameter.name == "video_path") - assert video_path_param.type_name == "str" + assert video_path_param.type_name == "Path | str" assert all(parameter.name != "output_path" for parameter in video_entry.parameters) - 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_types_preserve_all_registry_parameters(self, upload_service: ConverterService) -> None: + result = await upload_service.list_converter_types_async() + metadata_by_name = { + metadata.class_name: metadata for metadata in upload_service._registry.get_all_registered_class_metadata() + } + + assert {entry.converter_type for entry in result.items} == set(metadata_by_name) + for entry in result.items: + assert entry.parameters == list(metadata_by_name[entry.converter_type].parameters) + + @pytest.mark.parametrize( + ("converter_type", "parameter_name", "type_name", "required", "is_list"), + [ + ("SearchReplaceConverter", "replace", "str | list[str]", True, False), + ("DenylistConverter", "denylist", "list[str]", False, True), + ], + ) + async def test_types_expose_structured_parameters_without_changing_catalog( + self, + upload_service: ConverterService, + converter_type: str, + parameter_name: str, + type_name: str, + required: bool, + is_list: bool, + ) -> None: + types_result = await upload_service.list_converter_types_async() + catalog_result = await upload_service.list_converter_catalog_async() + types_entry = next(entry for entry in types_result.items if entry.converter_type == converter_type) + catalog_entry = next(entry for entry in catalog_result.items if entry.converter_type == converter_type) + parameter = next(param for param in types_entry.parameters if param.name == parameter_name) + + assert parameter.type_name == type_name + assert parameter.required is required + assert parameter.is_list is is_list + assert catalog_entry.parameters == [param for param in types_entry.parameters if param.is_string_coercible] + assert all(param.name != parameter_name for param in catalog_entry.parameters) + + 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" + + @pytest.mark.parametrize( + ("converter_type", "parameter_name"), + [ + ("AddImageTextConverter", "img_to_add"), + ("AddImageTextConverter", "font_name"), + ("AddTextImageConverter", "font_name"), + ("ColloquialWordswapConverter", "wordswap_path"), + ("ImagePromptStyleConverter", "filter_path"), + ("PDFConverter", "existing_pdf"), + ("TransparencyAttackConverter", "benign_image_path"), + ], + ) + async def test_local_constructor_files_use_path_parameters(self, converter_type: str, parameter_name: str) -> None: + service = ConverterService() + + result = await service.list_converter_types_async() + + entry = next(item for item in result.items if item.converter_type == converter_type) + parameter = next(item for item in entry.parameters if item.name == parameter_name) + assert parameter.is_path is True + + @pytest.mark.parametrize( + ("converter_type", "parameter_name"), + [("AddImageVideoConverter", "video_path"), ("ImageOverlayConverter", "base_image")], + ) + async def test_types_include_path_or_str_parameters(self, converter_type: str, parameter_name: str) -> None: + service = ConverterService() + + result = await service.list_converter_types_async() + + entry = next(item for item in result.items if item.converter_type == converter_type) + parameter = next(item for item in entry.parameters if item.name == parameter_name) + assert parameter.type_name == "Path | str" + assert parameter.is_path_or_str is True class TestGetConverter: @@ -248,6 +363,7 @@ async def test_create_converter_raises_for_invalid_type(self) -> None: service = ConverterService() request = CreateConverterRequest( + name="invalid", type="NonExistentConverter", params={}, ) @@ -260,22 +376,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={}, ) @@ -286,87 +403,419 @@ 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() + + result = await service.create_converter_async( + request=CreateConverterRequest(type="Base64Converter", params={}), + ) -class TestPersistDataUriParams: - """Tests for ConverterService._persist_data_uri_params_async (registry-metadata driven).""" + assert result.converter_id + assert service.get_converter_object(converter_id=result.converter_id) is not None - 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.""" + 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_preserves_replacement_registered_during_cleanup(self) -> None: + service = ConverterService() + original = Base64Converter() + replacement = Base64Converter() + service._registry.instances.register(original, name="converter") + + async def replace_during_cleanup_async(*, paths: list[Path]) -> None: + assert paths == [] + original_entry = service._registry.instances.get_entry("converter") + assert original_entry is not None + service._registry.instances.unregister("converter", expected_entry=original_entry) + service._registry.instances.register(replacement, name="converter") + + with patch.object(service, "_remove_owned_artifacts_async", side_effect=replace_during_cleanup_async): + removed = await service.delete_converter_async(converter_id="converter") + + assert removed is False + assert service._registry.instances.get("converter") is replacement + + async def test_delete_converter_removes_only_explicitly_owned_uploads( + self, upload_service: ConverterService + ) -> None: + service = upload_service + 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}) + + 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() + assert service._upload_path.is_dir() + + 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", + type_name="PDFConverter", + params={"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).""" + + @pytest.mark.parametrize( + ("converter_type", "parameter_name", "mime_type", "extension"), + [ + ("AddImageVideoConverter", "video_path", "video/mp4", ".mp4"), + ("ImageOverlayConverter", "base_image", "image/png", ".png"), + ], + ) + async def test_create_with_path_or_str_upload( + self, + upload_service: ConverterService, + converter_type: str, + parameter_name: str, + mime_type: str, + extension: str, + ) -> None: + content = b"uploaded content" + response = await upload_service.create_converter_async( + request=CreateConverterRequest( + name="uploaded", + type=converter_type, + params={parameter_name: _make_data_uri(mime_type=mime_type, content=content)}, + ) + ) + + entry = upload_service._registry.instances.get_entry(response.converter_id) + assert entry is not None + path = Path(entry.instance.get_identifier().params[parameter_name]) + assert path.parent == upload_service._upload_path + assert path.suffix == extension + assert path.read_bytes() == content + assert entry.metadata["owned_artifact_paths"] == [str(path)] + assert await upload_service.delete_converter_async(converter_id=response.converter_id) + assert not path.exists() + + @pytest.mark.parametrize( + ("converter_type", "parameter_name", "extension"), + [("AddImageVideoConverter", "video_path", "mp4"), ("ImageOverlayConverter", "base_image", "png")], + ) + async def test_create_with_path_or_str_url( + self, upload_service: ConverterService, converter_type: str, parameter_name: str, extension: str + ) -> None: + url = f"https://account.blob.core.windows.net/container/input.{extension}" + response = await upload_service.create_converter_async( + request=CreateConverterRequest(name="remote", type=converter_type, params={parameter_name: url}) + ) + + entry = upload_service._registry.instances.get_entry(response.converter_id) + assert entry is not None + assert entry.instance.get_identifier().params[parameter_name] == url + assert entry.metadata["owned_artifact_paths"] == [] + assert list(upload_service._upload_path.iterdir()) == [] + assert await upload_service.delete_converter_async(converter_id=response.converter_id) + + @pytest.mark.parametrize("value", [r"C:\server\input.mp4", "input.mp4", "https://example.org/input.mp4", 123]) + async def test_path_or_str_rest_rejects_non_upload_non_blob_values( + self, upload_service: ConverterService, value: object + ) -> None: + with pytest.raises(ValueError, match="data URI or supplied as an Azure Blob URL"): + await upload_service.create_converter_async( + request=CreateConverterRequest( + name="invalid", type="AddImageVideoConverter", params={"video_path": value} + ) + ) + assert upload_service._registry.instances.get_entry("invalid") is None + assert list(upload_service._upload_path.iterdir()) == [] + + async def test_plain_string_does_not_enable_upload_handling(self, upload_service: ConverterService) -> None: + value = _make_data_uri(mime_type="text/plain", content=b"literal suffix") + result, owned_paths = await upload_service._persist_data_uri_params_async( + converter_type="SuffixAppendConverter", params={"suffix": value} + ) + assert result == {"suffix": value} + assert owned_paths == [] + + async def test_persist_data_uri_materializes_path_in_managed_local_directory( + self, upload_service: ConverterService + ) -> None: + """A ``Path`` upload stays local even when CentralMemory storage is not local.""" + service = upload_service + memory = MagicMock(spec=MemoryInterface) + memory.results_path = "https://account.blob.core.windows.net/results" + params = {"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")} + + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory, ): - result = await service._persist_data_uri_params_async( - converter_type="ColloquialWordswapConverter", params=params + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="PDFConverter", + params=params, ) - assert result["wordswap_path"] == "/tmp/words.yaml" - assert not isinstance(result["wordswap_path"], Path) + assert result["existing_pdf"].is_absolute() + assert result["existing_pdf"].parent == service._upload_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_handles_optional_path_parameters(self) -> None: + service = ConverterService() + data_uri = _make_data_uri(mime_type="text/yaml", content=b"hello") + params = {"wordswap_path": data_uri} + + result, owned_paths = await service._persist_data_uri_params_async( + converter_type="ColloquialWordswapConverter", params=params + ) + assert result["wordswap_path"] == owned_paths[0] + assert owned_paths[0].read_bytes() == b"hello" 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, upload_service: ConverterService + ) -> None: + """Uploads are stored verbatim; restricting content is the media route's job.""" + service = upload_service + params = {"existing_pdf": _make_data_uri(mime_type=mime_type, content=b"")} + + 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, upload_service: ConverterService) -> None: + service = upload_service + params = {"existing_pdf": "data:application/pdf;base64,not-base64!!"} + + with pytest.raises(ValueError, match="invalid base64 data"): + await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + + assert list(service._upload_path.iterdir()) == [] + + async def test_persist_data_uri_rejects_non_base64_data_uri(self, upload_service: ConverterService) -> None: + service = upload_service + params = {"existing_pdf": "data:text/plain,hello"} + + with pytest.raises(ValueError, match="must be a base64 data URI"): + await service._persist_data_uri_params_async(converter_type="PDFConverter", params=params) + + assert list(service._upload_path.iterdir()) == [] + + async def test_create_converter_cleans_upload_when_construction_fails( + self, upload_service: ConverterService + ) -> None: + service = upload_service + params = { + "existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n"), + "font_color": [256, 0, 0], + } + request = CreateConverterRequest(name="invalid-pdf", type="PDFConverter", params=params) + + with 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(service._upload_path.iterdir()) == [] + + @pytest.mark.parametrize("error", [OSError("write failed"), asyncio.CancelledError()]) + async def test_persist_data_uri_cleans_partial_write( + self, upload_service: ConverterService, error: BaseException + ) -> None: + async def fail_write_async(content: bytes) -> None: + file_path = mock_open.call_args.args[0] + await asyncio.to_thread(file_path.write_bytes, content[:3]) + raise error + + request = CreateConverterRequest( + name="failed-upload", + type="PDFConverter", + params={"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")}, + ) + with patch("pyrit.backend.services.converter_service.aiofiles.open") as mock_open: + mock_file = mock_open.return_value.__aenter__.return_value + mock_file.write.side_effect = fail_write_async + with pytest.raises(type(error)): + await upload_service.create_converter_async(request=request) + + assert upload_service._registry.instances.get("failed-upload") is None + assert list(upload_service._upload_path.iterdir()) == [] + + async def test_concurrent_uploads_share_one_temporary_directory(self, upload_service: ConverterService) -> None: + params = {"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")} + results = await asyncio.gather( + upload_service._persist_data_uri_params_async(converter_type="PDFConverter", params=params), + upload_service._persist_data_uri_params_async(converter_type="PDFConverter", params=params), + ) + paths = [result["existing_pdf"] for result, _ in results] + assert paths[0] != paths[1] + assert all(path.parent == upload_service._upload_path for path in paths) + assert all(path.read_bytes() == b"%PDF-1.4\n" for path in paths) + + +class TestConverterServiceCleanup: + async def test_close_removes_only_owned_inputs(self, upload_service: ConverterService, tmp_path: Path) -> None: + service = upload_service + caller_file = tmp_path / "caller-owned.pdf" + caller_file.write_bytes(b"%PDF-1.4\n") + service._registry.create_named_instance( + name="caller-owned", + type_name="PDFConverter", + params={"existing_pdf": caller_file}, + ) + service._registry.create_named_instance(name="no-upload", type_name="Base64Converter") + request = CreateConverterRequest( + name="owned", + type="PDFConverter", + params={"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")}, + ) + await service.create_converter_async(request=request) + await service.close_async() + + assert not service._upload_path.exists() + assert service._registry.instances.get("owned") is None + assert service._registry.instances.get("caller-owned") is not None + assert service._registry.instances.get("no-upload") is not None + assert caller_file.read_bytes() == b"%PDF-1.4\n" + + async def test_close_keeps_other_service_uploads(self, upload_service: ConverterService) -> None: + other_service = ConverterService() + try: + assert upload_service._upload_path != other_service._upload_path + await other_service.create_converter_async( + request=CreateConverterRequest( + name="other", + type="PDFConverter", + params={"existing_pdf": _make_data_uri(mime_type="application/pdf", content=b"%PDF-1.4\n")}, + ) + ) + entry = other_service._registry.instances.get_entry("other") + assert entry is not None + owned_path = Path(entry.metadata["owned_artifact_paths"][0]) + + await upload_service.close_async() + + assert owned_path.read_bytes() == b"%PDF-1.4\n" + assert other_service._registry.instances.get("other") is entry.instance + finally: + await other_service.close_async() + + async def test_close_propagates_cleanup_errors(self, upload_service: ConverterService) -> None: + with patch.object(upload_service._upload_directory, "cleanup", side_effect=PermissionError("file in use")): + with pytest.raises(PermissionError, match="file in use"): + await upload_service.close_async() + + assert upload_service._upload_path.is_dir() + class TestPreviewConversion: """Tests for ConverterService.preview_conversion method.""" diff --git a/tests/unit/backend/test_main.py b/tests/unit/backend/test_main.py index d91c59d0f6..6a8f883a5e 100644 --- a/tests/unit/backend/test_main.py +++ b/tests/unit/backend/test_main.py @@ -9,6 +9,7 @@ import logging import os +from contextlib import nullcontext from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch @@ -18,12 +19,61 @@ from starlette.exceptions import HTTPException as StarletteHTTPException from pyrit.backend.main import SPAStaticFiles, app, lifespan, setup_frontend +from pyrit.backend.models.converters import CreateConverterRequest +from pyrit.backend.services.converter_service import get_converter_service from pyrit.setup.configuration_loader import ConfigurationLoader class TestLifespan: """Tests for the application lifespan context manager.""" + @pytest.mark.parametrize("fail_during_lifespan", [False, True]) + async def test_lifespan_cleans_converter_uploads(self, fail_during_lifespan: bool) -> None: + fake_config = ConfigurationLoader() + with ( + patch.object(ConfigurationLoader, "load_with_overrides", return_value=fake_config), + patch.object(ConfigurationLoader, "initialize_pyrit_async", new=AsyncMock()), + patch("pyrit.backend.main.setup_frontend"), + pytest.raises(RuntimeError, match="application failed") if fail_during_lifespan else nullcontext(), + ): + async with lifespan(app): + service = get_converter_service() + await service.create_converter_async( + request=CreateConverterRequest( + name="lifespan-upload", + type="PDFConverter", + params={"existing_pdf": "data:application/pdf;base64,JVBERi0xLjQK"}, + ) + ) + entry = service._registry.instances.get_entry("lifespan-upload") + assert entry is not None + owned_path = Path(entry.metadata["owned_artifact_paths"][0]) + assert owned_path.read_bytes() == b"%PDF-1.4\n" + if fail_during_lifespan: + raise RuntimeError("application failed") + + assert not owned_path.exists() + assert not service._upload_path.exists() + assert service._registry.instances.get("lifespan-upload") is None + assert get_converter_service.cache_info().currsize == 0 + + async def test_lifespan_restarts_with_fresh_upload_directory(self) -> None: + fake_config = ConfigurationLoader() + paths: list[Path] = [] + with ( + patch.object(ConfigurationLoader, "load_with_overrides", return_value=fake_config), + patch.object(ConfigurationLoader, "initialize_pyrit_async", new=AsyncMock()), + patch("pyrit.backend.main.setup_frontend"), + ): + for _ in range(2): + async with lifespan(app): + path = get_converter_service()._upload_path + assert path.is_dir() + paths.append(path) + assert not path.exists() + + assert paths[0] != paths[1] + async def test_lifespan_yields(self) -> None: """Test that lifespan delegates to ConfigurationLoader and yields.""" fake_config = ConfigurationLoader() diff --git a/tests/unit/backend/test_mappers.py b/tests/unit/backend/test_mappers.py index 0d75a3657b..8d8924bb63 100644 --- a/tests/unit/backend/test_mappers.py +++ b/tests/unit/backend/test_mappers.py @@ -1946,13 +1946,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.""" @@ -1963,7 +1969,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..efb52ede86 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: @@ -126,14 +127,36 @@ def test_serves_file_from_seed_prompt_entries(self, client: TestClient, _mock_me assert response.status_code == 200 - def test_rejects_unknown_extension(self, client: TestClient, _mock_memory: Path) -> None: - """Files with unknown extensions are rejected by the allowlist.""" - file_path = _mock_memory / "prompt-memory-entries" / "data.xyz123" - file_path.write_bytes(b"binary data") + @pytest.mark.parametrize( + ("file_name", "content"), + [ + ("program.exe", b"MZ"), + ("data.xyz123", b"binary data"), + ("config.yaml", b"key: value"), + ("leaked.db", b"SQLite format 3"), + ("active.html", b""), + ("active.svg", b""), + ], + ) + def test_non_inline_type_downloads_as_opaque_attachment( + self, + client: TestClient, + _mock_memory: Path, + file_name: str, + content: bytes, + ) -> None: + """Any stored type can download, but only allowlisted media renders inline.""" + file_path = _mock_memory / "prompt-memory-entries" / file_name + file_path.write_bytes(content) response = client.get("/api/media", params={"path": str(file_path)}) - assert response.status_code == 403 + assert response.status_code == 200 + assert response.content == content + assert response.headers["content-type"] == "application/octet-stream" + assert response.headers["content-disposition"].startswith("attachment;") + assert file_name in response.headers["content-disposition"] + assert response.headers["x-content-type-options"] == "nosniff" def test_rejects_file_in_results_root(self, client: TestClient, _mock_memory: Path) -> None: """Files directly in results_path (not in allowed subdir) are rejected.""" @@ -144,23 +167,17 @@ def test_rejects_file_in_results_root(self, client: TestClient, _mock_memory: Pa assert response.status_code == 403 - def test_rejects_database_file_in_allowed_subdir(self, client: TestClient, _mock_memory: Path) -> None: - """Database files are not in the extension allowlist.""" - file_path = _mock_memory / "prompt-memory-entries" / "leaked.db" - file_path.write_bytes(b"SQLite format 3") - - response = client.get("/api/media", params={"path": str(file_path)}) - - assert response.status_code == 403 - - def test_rejects_yaml_file(self, client: TestClient, _mock_memory: Path) -> None: - """YAML files are not in the extension allowlist.""" - file_path = _mock_memory / "prompt-memory-entries" / "config.yaml" - file_path.write_bytes(b"key: value") + def test_serves_documents_as_attachments(self, client: TestClient, _mock_memory: Path) -> None: + """Documents download as opaque bytes 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 == 403 + assert response.status_code == 200 + assert response.headers["content-type"] == "application/octet-stream" + 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.""" diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index c4ff0f41c2..94c6768f3c 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -257,6 +257,37 @@ 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 targets_parameter.type_name == "list[str]" + assert targets_parameter.is_list is True + weights_parameter = next(param for param in types_entry.parameters if param.name == "weights") + assert weights_parameter.type_name == "list[int]" + assert weights_parameter.is_list is True + assert weights_parameter.required is False + assert all(param.name != "targets" for param in catalog_entry.parameters) + assert all(param.name != "weights" 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_preserve_all_registry_parameters(self) -> None: + service = TargetService() + result = await service.list_target_types_async() + metadata_by_name = { + metadata.class_name: metadata for metadata in service._registry.get_all_registered_class_metadata() + } + + assert {entry.target_type for entry in result.items} == set(metadata_by_name) + for entry in result.items: + assert entry.parameters == list(metadata_by_name[entry.target_type].parameters) + async def test_catalog_cold_and_warm_results_are_equal(self) -> None: service = TargetService() @@ -361,6 +392,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 73% rename from tests/unit/backend/test_target_catalog_concurrency.py rename to tests/unit/backend/test_target_types_concurrency.py index 9599a009f3..59b6e83e7f 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,15 +31,15 @@ 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) try: - health_response = await client.get("/api/health") + health_response = await asyncio.wait_for(client.get("/api/health"), timeout=2) assert health_response.status_code == 200 assert not discovery_finished.is_set() finally: discovery_release.set() - catalog_response = await catalog_request + 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..97c61e7377 100644 --- a/tests/unit/models/test_parameter.py +++ b/tests/unit/models/test_parameter.py @@ -4,7 +4,8 @@ """Unit tests for the unified Parameter model and its coercion methods.""" from enum import Enum -from typing import Literal +from pathlib import Path +from typing import Literal, Union import pytest from pydantic import ValidationError @@ -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,35 @@ 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 dumped["type_name"] == "str" + assert dumped["is_list"] is False + assert restored.reference == RegistryReference(component_type=ComponentType.TARGET, annotation=str) + assert restored.reference_type == "target" + + def test_list_reference_shape_round_trips(self) -> None: + parameter = Parameter( + name="targets", + description="d", + reference=RegistryReference(component_type=ComponentType.TARGET, annotation=list[object]), + ) + + dumped = parameter.model_dump() + restored = Parameter.model_validate(dumped) + + assert dumped["type_name"] == "list[str]" + assert dumped["is_list"] is True + assert restored.reference == RegistryReference(component_type=ComponentType.TARGET, annotation=list[str]) + 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 +165,73 @@ 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") + + def test_optional_path_is_path(self) -> None: + parameter = Parameter(name="input_path", description="d", param_type=Path | None) + + assert parameter.is_path is True + + +@pytest.mark.parametrize( + "annotation", + [Path | str, str | Path, Path | str | None, str | Path | None, Union[str, Path]], # noqa: UP007 +) +def test_path_or_str_contract_round_trip(annotation: object) -> None: + parameter = Parameter(name="source", description="d", param_type=annotation) + restored = Parameter.model_validate_json(parameter.model_dump_json()) + url = "https://account.blob.core.windows.net/container/input.png?versionid=123" + path = Path("input.png") + + for candidate in (parameter, restored): + candidate.validate() + assert candidate.is_path is False + assert candidate.is_path_or_str is True + assert candidate.is_string_coercible is True + assert candidate.type_name == "Path | str" + assert candidate.is_list is False + assert candidate.coerce_value(url) == url + assert candidate.coerce_value(path) is path + assert candidate.coerce_value("input.png") == "input.png" + with pytest.raises(ValueError, match="expects a Path or str"): + candidate.coerce_value(123) + + +@pytest.mark.parametrize("annotation", [str, Path, Path | int, str | int, Path | str | int, list[Path | str]]) +def test_path_or_str_does_not_match_other_types(annotation: object) -> None: + parameter = Parameter(name="source", description="d", param_type=annotation) + assert parameter.is_path_or_str is False + + +def test_optional_path_or_str_accepts_none() -> None: + parameter = Parameter(name="source", description="d", param_type=Path | str | None) + assert parameter.coerce_value(None) is None + + +def test_list_path_or_str_contract_round_trip() -> None: + parameter = Parameter(name="sources", description="d", param_type=list[Path | str]) + restored = Parameter.model_validate_json(parameter.model_dump_json()) + values = [Path("input.png"), "https://account.blob.core.windows.net/container/input.png"] + + for candidate in (parameter, restored): + candidate.validate() + assert candidate.type_name == "list[Path | str]" + assert candidate.is_list is True + assert candidate.is_string_coercible is False + assert candidate.coerce_value(values) == values + 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 +266,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 +333,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 +460,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..5a77897585 100644 --- a/tests/unit/registry/test_converter_registry.py +++ b/tests/unit/registry/test_converter_registry.py @@ -6,6 +6,7 @@ and its introspection helpers. """ +from pathlib import Path from typing import Literal import pytest @@ -116,6 +117,21 @@ def registry(): # --------------------------------------------------------------------------- +@pytest.mark.parametrize( + ("converter_type", "parameter_name"), + [("AddImageVideoConverter", "video_path"), ("ImageOverlayConverter", "base_image")], +) +@pytest.mark.parametrize( + "source", + [Path("input.png"), "input.png", "https://account.blob.core.windows.net/container/input.png"], +) +def test_registry_preserves_path_or_str_inputs( + registry: ConverterRegistry, converter_type: str, parameter_name: str, source: Path | str +) -> None: + instance = registry.create_instance(converter_type, **{parameter_name: source}) + assert instance.get_identifier().params[parameter_name] == str(source) + + class TestConverterRegistrySingleton: """Tests for the singleton pattern in ConverterRegistry.""" @@ -161,15 +177,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", type_name="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", + type_name="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, type_name="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..9eefe24b5b 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,43 @@ 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_expected_entry_preserves_replacement( + self, registry: DefaultInstanceRegistry[_TestItem] + ) -> None: + item = _item("item") + registry.register(item, name="name", metadata={"version": 1}) + original_entry = registry.get_entry("name") + assert original_entry is not None + registry.register(item, name="name", metadata={"version": 2}, replace=True) + + assert registry.unregister("name", expected_entry=original_entry) is None + assert registry.get("name") is item + replacement_entry = registry.get_entry("name") + assert replacement_entry is not None + assert replacement_entry.metadata == {"version": 2} + + 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..8ee2ffacf8 100644 --- a/tests/unit/registry/test_scorer_registry.py +++ b/tests/unit/registry/test_scorer_registry.py @@ -172,15 +172,31 @@ 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_create_named_instance_builds_and_stores_scorer(self, registry: ScorerRegistry): + registry.instances.register(MockTrueFalseScorer(), name="inner") + + scorer = registry.create_named_instance( + name="composite", + type_name="TrueFalseCompositeScorer", + params={ + "scorers": ["inner"], + "aggregator": TrueFalseScoreAggregator.OR, + }, + ) + + assert isinstance(scorer, TrueFalseCompositeScorer) + assert registry.instances.get("composite") is scorer 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..900ac2647d 100644 --- a/tests/unit/registry/test_target_registry.py +++ b/tests/unit/registry/test_target_registry.py @@ -124,15 +124,36 @@ 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", + type_name="MockPromptTarget", + params={"model_name": "named-model"}, + ) + + assert isinstance(target, MockPromptTarget) + assert target.get_identifier().model_name == "named-model" + 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, type_name="MockPromptTarget") def test_register_instance_rejects_non_target(self, registry: TargetRegistry): class NotATarget: diff --git a/tests/unit/scenario/airt/test_rapid_response.py b/tests/unit/scenario/airt/test_rapid_response.py index c32c63f483..1f40f58234 100644 --- a/tests/unit/scenario/airt/test_rapid_response.py +++ b/tests/unit/scenario/airt/test_rapid_response.py @@ -615,6 +615,7 @@ def test_register_from_factories_idempotent(self): def test_register_preserves_custom_preregistered(self): """Pre-registered custom techniques are not overwritten by re-registration.""" + AttackTechniqueRegistry.reset_registry_singleton() registry = AttackTechniqueRegistry.get_registry_singleton() custom_factory = AttackTechniqueFactory(name="role_play_movie_script", attack_class=PromptSendingAttack) registry.register_technique(name="role_play_movie_script", factory=custom_factory, tags=["custom"]) diff --git a/tests/unit/setup/test_converter_initializer.py b/tests/unit/setup/test_converter_initializer.py index b5ce7f3b49..f919d7a298 100644 --- a/tests/unit/setup/test_converter_initializer.py +++ b/tests/unit/setup/test_converter_initializer.py @@ -61,6 +61,25 @@ async def test_initialize_registers_variation_with_declared_target() -> None: assert converter._converter_target is adversarial_chat +@pytest.mark.usefixtures("patch_central_database") +async def test_initialize_replaces_converter_with_current_target() -> None: + converter_registry = ConverterRegistry.get_registry_singleton() + target_registry = TargetRegistry.get_registry_singleton() + first_target = MockPromptTarget() + target_registry.instances.register(first_target, name="adversarial_chat") + + with patch.object(ConverterInitializer, "CONFIGS", _get_configs("variation")): + await ConverterInitializer().initialize_async() + target_registry.instances.register(MockPromptTarget(), name="adversarial_chat", replace=True) + current_target = target_registry.instances.get("adversarial_chat") + await ConverterInitializer().initialize_async() + + converter = converter_registry.instances.get("variation") + assert isinstance(converter, VariationConverter) + assert converter._converter_target is current_target + assert converter._converter_target is not first_target + + async def test_initialize_skips_converter_without_declared_target() -> None: registry = ConverterRegistry.get_registry_singleton() 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