From 41b55b8d5df5e0e52fd3ab074d0e7df8bfe6e536 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 1 Oct 2026 19:43:55 +0000 Subject: [PATCH 1/7] FIX: Keep backend media and upload inputs inside managed storage Media values sent to the backend (message pieces, converter previews, and prepended conversations) could name any existing local file or any URL, and converter file parameters could name any Azure Blob URL. The server then read or forwarded those values. They must now be uploaded content or point into this server's media storage: the media folders under the memory results path, or, when results are stored in Azure Blob Storage, those folders in the configured results container. Blob URLs must start with the results URL and a media folder, the way the serializer builds them, both as the storage reader reads them and as HTTP clients decode them. Prepended content is checked before attack or conversation rows are written. Converter previews and sent messages now reject url pieces, so the backend no longer downloads URLs outside its result storage. Stored history can still hold them. Targets that read local files or load model code now declare it with uses_host_resources, and the create-target API rejects those types (HTTPXAPITarget, HuggingFaceChatTarget). They can still be registered in Python or with an initializer, where the operator controls their settings. The media route and media persistence now share one containment helper, and the backend README documents the checks and the intentional exceptions. --- pyrit/backend/README.md | 30 +++ pyrit/backend/models/attacks.py | 18 ++ pyrit/backend/models/converters.py | 18 +- pyrit/backend/routes/attacks.py | 7 +- pyrit/backend/routes/media.py | 25 +- pyrit/backend/services/attack_service.py | 28 ++- pyrit/backend/services/converter_service.py | 11 +- pyrit/backend/services/media_persistence.py | 216 ++++++++++++++++-- .../backend/services/message_send_service.py | 51 +---- pyrit/backend/services/target_service.py | 30 ++- pyrit/prompt_target/common/prompt_target.py | 5 + .../http_target/httpx_api_target.py | 4 +- .../hugging_face/hugging_face_chat_target.py | 5 +- tests/unit/backend/test_api_routes.py | 14 ++ tests/unit/backend/test_attack_models.py | 24 +- tests/unit/backend/test_attack_service.py | 58 +++++ tests/unit/backend/test_converter_service.py | 116 +++++++--- tests/unit/backend/test_media_persistence.py | 202 ++++++++++++++-- .../unit/backend/test_message_send_service.py | 142 +++++++++--- tests/unit/backend/test_target_service.py | 21 ++ 20 files changed, 821 insertions(+), 204 deletions(-) diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index c9e0c3420d..456f58db70 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -126,3 +126,33 @@ Environment variables: - `PYRIT_API_HOST` - Host to bind to (default: localhost) - `PYRIT_API_PORT` - Port to listen on (default: 8000) - `PYRIT_API_RELOAD` - Enable auto-reload (default: false) + +## Input Validation + +The backend is the part of PyRIT that accepts requests from other machines, so it checks +request values before using them: + +- Media values in messages, previews, and prepended conversations, and file parameters of + converters, must be uploaded content or point into this server's media storage: the + `prompt-memory-entries` and `seed-prompt-entries` folders under the memory results path. + Blob URLs are accepted only when results are stored in Azure Blob Storage, and only + inside those folders of the configured results container. Other file paths and URLs + are rejected. +- The backend only reads media URLs inside this server's result storage. Converter + previews and sent messages reject `url` pieces; stored history may still contain them. +- Target types that read local files or load model code (`HTTPXAPITarget`, + `HuggingFaceChatTarget`) cannot be created through the API. Register them in Python or + with an initializer, where the operator controls their settings; for example, set + `HTTPXAPITarget(allowed_upload_directory=...)` so uploads stay inside one folder. + +Intentional exceptions: + +- Target endpoints, raw HTTP requests, and their redirects are chosen by the operator and + are not restricted. Limit outbound network access in the deployment instead. +- Prompt content is not filtered. It is adversarial test data by design. +- Any file type can be stored as a payload. `GET /api/media` only renders known image, + audio, and video types inline; everything else downloads as a file. +- `GET /api/media` does not require authentication so the browser can load media. It only + serves files from the media folders above. +- Custom initializer scripts are trusted Python. Uploading them requires an administrator + and `allow_custom_initializers: true`. diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index 0cd5a99869..1419354ba6 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -667,6 +667,24 @@ def _validate_converter_configurations(self) -> "AddMessageRequest": return self + @model_validator(mode="after") + def _reject_url_pieces_when_sending(self) -> "AddMessageRequest": + """ + Reject URL pieces in a message that will be sent. + + Sending can route a piece through converters that download URL input on the + server. Stored history may still hold URL pieces because storing never fetches them. + + Returns: + AddMessageRequest: The validated request. + + Raises: + ValueError: If a piece to be sent has a ``url`` data type. + """ + if self.send and any("url" in (piece.data_type, piece.converted_value_data_type) for piece in self.pieces): + raise ValueError("URL pieces cannot be sent; upload the file content instead") + return self + class AddMessageResponse(BaseModel): """ diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 316cedb42b..66a37c5ec6 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -9,7 +9,7 @@ from typing import Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN from pyrit.models import ConverterIdentifier, Parameter, PromptDataType @@ -119,6 +119,22 @@ class ConverterPreviewRequest(BaseModel): original_value_data_type: PromptDataType = Field(default="text", description="Data type of original value") converter_ids: list[str] = Field(..., description="Converter instance IDs to apply") + @field_validator("original_value_data_type") + @classmethod + def _reject_url_input(cls, value: PromptDataType) -> PromptDataType: + """ + Reject URL input so a preview never makes the server download a caller-supplied URL. + + Returns: + PromptDataType: The validated data type. + + Raises: + ValueError: If the data type is ``url``. + """ + if value == "url": + raise ValueError("URL input is not supported; upload the file content instead") + return value + class ConverterPreviewResponse(BaseModel): """Response from converter preview.""" diff --git a/pyrit/backend/routes/attacks.py b/pyrit/backend/routes/attacks.py index f5935f4e43..53a90a1d8b 100644 --- a/pyrit/backend/routes/attacks.py +++ b/pyrit/backend/routes/attacks.py @@ -219,10 +219,9 @@ async def create_attack(request: CreateAttackRequest) -> CreateAttackResponse: try: return await service.create_attack_async(request=request) except ValueError as e: - raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail=str(e), - ) from e + error_msg = str(e) + error_status = status.HTTP_404_NOT_FOUND if "not found" in error_msg.lower() else status.HTTP_400_BAD_REQUEST + raise HTTPException(status_code=error_status, detail=error_msg) from e @router.get( diff --git a/pyrit/backend/routes/media.py b/pyrit/backend/routes/media.py index da93d1e3e1..116ba8ae36 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -23,15 +23,13 @@ from fastapi import APIRouter, HTTPException, Query from fastapi.responses import FileResponse +from pyrit.backend.services.media_persistence import resolve_managed_media_path from pyrit.memory import CentralMemory logger = logging.getLogger(__name__) router = APIRouter() -# Only serve files from known media subdirectories under results_path. -_ALLOWED_SUBDIRECTORIES = {"prompt-memory-entries", "seed-prompt-entries"} - # Only these known-safe media types render inline. Every other extension is # served as an application/octet-stream attachment. _INLINE_EXTENSIONS = { @@ -62,12 +60,7 @@ def _validate_media_path(*, path: str, allowed_root: Path) -> Path: """ - Validate and sanitize a user-provided file path against an allowed root directory. - - Uses ``Path.resolve()`` to resolve symlinks and ``..`` components, then - verifies the canonical path is under the allowed root. This is the standard - sanitization pattern recognized by static analysis tools (e.g. CodeQL - ``py/path-injection``). + Validate a user-provided file path against the allowed results directory. Args: path: The user-provided file path to validate. @@ -79,20 +72,10 @@ def _validate_media_path(*, path: str, allowed_root: Path) -> Path: Raises: HTTPException 403: If the path fails any validation check. """ - real_path = Path(path).resolve(strict=False) - try: - relative_parts = real_path.relative_to(allowed_root).parts + return resolve_managed_media_path(path=path, allowed_root=allowed_root) except ValueError as exc: - raise HTTPException( - status_code=403, detail="Access denied: path is outside the allowed results directory." - ) from exc - - # Restrict to known media subdirectories (e.g. prompt-memory-entries/) - 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.") - - return real_path + raise HTTPException(status_code=403, detail=f"Access denied: {exc}") from exc @router.get("/media") diff --git a/pyrit/backend/services/attack_service.py b/pyrit/backend/services/attack_service.py index 2527c7b2fb..b94e9a500c 100644 --- a/pyrit/backend/services/attack_service.py +++ b/pyrit/backend/services/attack_service.py @@ -48,6 +48,7 @@ UpdateMainConversationResponse, ) from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.services.media_persistence import persist_message_pieces_async from pyrit.backend.services.message_send_service import MessageSendService, resolve_applied_converter_identifiers from pyrit.backend.services.pagination import ( decode_keyset_cursor, @@ -57,7 +58,7 @@ ) from pyrit.backend.services.target_service import get_target_service from pyrit.common.utils import to_sha256 -from pyrit.memory import AttackResultKeysetCursor, CentralMemory +from pyrit.memory import AttackResultKeysetCursor, CentralMemory, data_serializer_factory from pyrit.models import ( AtomicAttackIdentifier, AttackIdentifier, @@ -390,6 +391,20 @@ async def create_attack_async(self, *, request: CreateAttackRequest) -> CreateAt labels.setdefault("source", "gui") attack_result_id = str(uuid.uuid4()) + # A system_prompt is lowered to a single system-role message at the front, composing + # with any prepended_conversation. Media is checked before attack or conversation rows are written. + prepended = list(request.prepended_conversation or []) + if request.system_prompt: + prepended.insert( + 0, + PrependedMessageRequest( + role="system", + pieces=[MessagePieceRequest(original_value=request.system_prompt)], + ), + ) + for message in prepended: + await persist_message_pieces_async(pieces=message.pieces, serializer_factory=data_serializer_factory) + # --- Branch via duplication (preferred for tracking) --------------- if request.source_conversation_id is not None and request.cutoff_index is not None: conversation_id = await self._duplicate_conversation_up_to_async( @@ -439,17 +454,6 @@ async def create_attack_async(self, *, request: CreateAttackRequest) -> CreateAt ) (await self._memory.add_attack_results_to_memory_async(attack_results=[attack_result])) - # Store prepended conversation messages if provided. A system_prompt is lowered to a - # single system-role message at the front, composing with any prepended_conversation. - prepended = list(request.prepended_conversation or []) - if request.system_prompt: - prepended.insert( - 0, - PrependedMessageRequest( - role="system", - pieces=[MessagePieceRequest(original_value=request.system_prompt)], - ), - ) if prepended: await self._store_prepended_messages_async( conversation_id=conversation_id, diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index e85f91be47..988970efb0 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -37,7 +37,7 @@ CreateConverterRequest, PreviewStep, ) -from pyrit.backend.services.media_persistence import persist_media_value_async +from pyrit.backend.services.media_persistence import persist_media_value_async, require_managed_blob_url from pyrit.common.azure_storage import is_azure_blob_uri from pyrit.memory import data_serializer_factory from pyrit.models import MessagePiece, PromptDataType @@ -284,8 +284,9 @@ async def _persist_data_uri_params_async( 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. + ``Path | str`` parameters also accept Azure Blob URLs inside this server's + result storage, 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. @@ -303,7 +304,8 @@ async def _persist_data_uri_params_async( set of request-created files owned by the future registry entry. Raises: - ValueError: If a ``Path`` value is not a valid data URI. + ValueError: If a ``Path`` value is not a valid data URI, or a blob URL + points outside this server's result storage. """ metadata = self._registry.get_registered_class_metadata(converter_type) path_params = ( @@ -327,6 +329,7 @@ async def _persist_data_uri_params_async( 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): + require_managed_blob_url(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}") diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index e9c73d6989..2cbb024d16 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -5,22 +5,29 @@ from __future__ import annotations +import asyncio import base64 import binascii import mimetypes -from collections.abc import Callable +from collections.abc import Callable, Sequence from dataclasses import dataclass from enum import Enum from pathlib import Path from typing import TYPE_CHECKING, Any -from urllib.parse import parse_qs, urlparse +from urllib.parse import parse_qs, unquote, urlparse from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS -from pyrit.memory import data_serializer_factory +from pyrit.common.azure_storage import is_azure_blob_uri +from pyrit.memory import CentralMemory, data_serializer_factory +from pyrit.models import MEDIA_PATH_DATA_TYPES if TYPE_CHECKING: + from pyrit.backend.models.attacks import MessagePieceRequest from pyrit.models import PromptDataType +# Media is only read from these folders under the memory results path. +MEDIA_SUBDIRECTORIES = frozenset({"prompt-memory-entries", "seed-prompt-entries"}) + class MediaOrigin(str, Enum): """Origin recognized for one path-typed media value.""" @@ -63,6 +70,118 @@ def _data_uri_parts(value: str) -> tuple[str | None, str]: return media_type, payload +def resolve_managed_media_path(*, path: str, allowed_root: Path) -> Path: + """ + Return the canonical form of a local media path inside the results directory. + + ``Path.resolve()`` removes ``..`` parts and follows symlinks before the + containment check, so the returned path is the file that will actually be read. + + Args: + path (str): The caller-supplied file path. + allowed_root (Path): The canonical (``resolve``-d) memory results directory. + + Returns: + Path: The canonical path. + + Raises: + ValueError: If the path is outside ``allowed_root`` or not in a media folder. + """ + real_path = Path(path).resolve(strict=False) + try: + relative_parts = real_path.relative_to(allowed_root).parts + except ValueError as exc: + raise ValueError("Media path is outside the allowed results directory.") from exc + if not relative_parts or relative_parts[0] not in MEDIA_SUBDIRECTORIES: + raise ValueError("Media path is not in a media folder of the allowed results directory.") + return real_path + + +def _results_root() -> str: + """Return the memory results path, which is a local folder or a blob container URL.""" + return str(CentralMemory.get_memory_instance().results_path) + + +def _is_remote_root(root: str) -> bool: + """Return whether the results path is a blob container URL rather than a local folder.""" + return urlparse(root).scheme in ("http", "https") + + +def _resolve_local_media_value(value: str) -> Path: + """ + Resolve a local media path against the configured results folder. + + Returns: + Path: The canonical media path. + + Raises: + ValueError: If results are stored remotely or the path is outside managed media. + """ + root = _results_root() + if _is_remote_root(root): + raise ValueError("Local media paths are not accepted when results are stored in Azure Blob Storage.") + return resolve_managed_media_path(path=value, allowed_root=Path(root).resolve(strict=False)) + + +def _is_media_blob_path(*, path: str, prefix: str) -> bool: + """ + Return whether a URL path names a blob inside a media folder of the results container. + + Returns: + bool: True when the path is the results prefix, a media folder, and a blob name without empty or dot + segments. + """ + path = path.replace("\\", "/") + if not path.startswith(prefix): + return False + blob_segments = path[len(prefix) :].split("/") + return ( + len(blob_segments) > 1 and blob_segments[0] in MEDIA_SUBDIRECTORIES and not {"", ".", ".."} & set(blob_segments) + ) + + +def require_managed_blob_url(value: str) -> None: + """ + Require a media URL to point into a media folder of the configured results container. + + Raises: + ValueError: If the URL is not an Azure Blob URL inside managed result storage. + """ + root = _results_root() + if not _is_remote_root(root) or not is_azure_blob_uri(value): + raise ValueError("Media URLs must point to this server's result storage.") + container, candidate = urlparse(root), urlparse(value.replace("\\", "/")) + # Media URLs are the results URL followed by a media folder, the way the serializer builds them. + # The blob storage reader uses the path without percent-decoding it, while HTTP clients decode + # it first, so both readings must stay inside a media folder. + prefix = f"{container.path}/" + if ( + not container.path.strip("/") + or candidate.netloc.lower() != container.netloc.lower() + or not all(_is_media_blob_path(path=path, prefix=prefix) for path in (candidate.path, unquote(candidate.path))) + ): + raise ValueError("Media URLs must point to this server's result storage.") + + +def _media_reference_path(value: str) -> str | None: + """ + Return the file path from an ``/api/media?path=...`` reference. + + Returns: + str | None: The referenced path, or None when the value is not a media reference. + + Raises: + ValueError: If the reference does not carry exactly one non-empty ``path`` value. + """ + parsed = urlparse(value) + if parsed.scheme or parsed.netloc or parsed.path != "/api/media": + return None + paths = parse_qs(parsed.query).get("path", []) + if len(paths) != 1 or not paths[0]: + raise ValueError("Media references must include exactly one path.") + return paths[0] + + def _resolve_extension( *, data_type: PromptDataType, @@ -96,12 +215,18 @@ async def persist_media_value_async( The two policy flags preserve the small historical differences between attack ingestion and converter preview while keeping origin detection, - extension resolution, and persistence in one component. + extension resolution, and persistence in one component. Existing files, + ``/api/media`` references, and URLs are accepted only when they point into + this server's media storage, so callers cannot make the server read other files. Returns: A typed result containing the resolved value and persistence metadata. + + Raises: + ValueError: If the value names a file or URL outside the server's media storage. """ if value.startswith(("http://", "https://")): + require_managed_blob_url(value) return MediaPersistenceResult( value=value, origin=MediaOrigin.REMOTE_URL, @@ -110,14 +235,14 @@ async def persist_media_value_async( mime_type=mime_type, ) - if value.startswith("/api/media"): - parsed = urlparse(value) - file_path = parse_qs(parsed.query).get("path", [None])[0] + reference_path = _media_reference_path(value) + if reference_path is not None: + managed_path = await asyncio.to_thread(_resolve_local_media_value, reference_path) return MediaPersistenceResult( - value=file_path or value, + value=str(managed_path), origin=MediaOrigin.MEDIA_REFERENCE, persisted=False, - resolved=file_path is not None, + resolved=True, mime_type=mime_type, ) @@ -129,17 +254,20 @@ async def persist_media_value_async( origin = MediaOrigin.DATA_URI else: try: - if Path(value).is_file(): - return MediaPersistenceResult( - value=value, - origin=MediaOrigin.LOCAL_PATH, - persisted=False, - resolved=True, - mime_type=mime_type, - ) + is_local_file = await asyncio.to_thread(Path(value).is_file) except (OSError, ValueError): if require_valid_base64_after_path_error and not _is_raw_base64(value): raise + is_local_file = False + if is_local_file: + managed_path = await asyncio.to_thread(_resolve_local_media_value, value) + return MediaPersistenceResult( + value=str(managed_path), + origin=MediaOrigin.LOCAL_PATH, + persisted=False, + resolved=True, + mime_type=mime_type, + ) extension = _resolve_extension( data_type=data_type, @@ -161,3 +289,57 @@ async def persist_media_value_async( mime_type=mime_type or data_uri_mime_type, extension=extension, ) + + +async def persist_message_pieces_async( + *, + pieces: Sequence[MessagePieceRequest], + serializer_factory: SerializerFactory = data_serializer_factory, +) -> None: + """ + Resolve original and converted media independently, updating values in-place. + + The frontend sends binary media (images, audio, etc.) as base64 strings + with a ``*_path`` data_type. The PyRIT target layer expects ``*_path`` + values to be **file paths**, so base64 data is written to the results + store and the request values are replaced with the resulting file path. + Values that already reference stored media are kept after they are + checked to be inside this server's media storage. + + Args: + pieces (Sequence[MessagePieceRequest]): Request pieces to resolve in place. + serializer_factory (SerializerFactory): Factory used to write new media files. + """ + for piece in pieces: + original_value = piece.original_value + converted_value = piece.converted_value + converted_type = piece.converted_value_data_type or piece.data_type + if piece.data_type in MEDIA_PATH_DATA_TYPES: + result = await persist_media_value_async( + value=original_value, + data_type=piece.data_type, + mime_type=piece.mime_type, + serializer_factory=serializer_factory, + ) + if result.resolved: + original_value = result.value + if converted_value is None or ( + converted_value == piece.original_value and converted_type == piece.data_type + ): + converted_value = original_value + + if ( + converted_value is not None + and converted_type in MEDIA_PATH_DATA_TYPES + and (converted_value != original_value or converted_type != piece.data_type) + ): + result = await persist_media_value_async( + value=converted_value, + data_type=converted_type, + serializer_factory=serializer_factory, + ) + if result.resolved: + converted_value = result.value + + piece.original_value = original_value + piece.converted_value = converted_value diff --git a/pyrit/backend/services/message_send_service.py b/pyrit/backend/services/message_send_service.py index 11a0915626..36bb808fa3 100644 --- a/pyrit/backend/services/message_send_service.py +++ b/pyrit/backend/services/message_send_service.py @@ -15,13 +15,12 @@ from pyrit.backend.models.attacks import AddMessageRequest, ConverterConfigurationRequest, MessagePieceRequest from pyrit.backend.services.converter_service import get_converter_service from pyrit.backend.services.manual_send_scheduler import ManualSendScheduler, get_manual_send_scheduler -from pyrit.backend.services.media_persistence import persist_media_value_async +from pyrit.backend.services.media_persistence import persist_message_pieces_async from pyrit.backend.services.target_service import get_target_service from pyrit.common.attack_result_scope import attack_result_id_scope from pyrit.common.deprecation import print_deprecation_message from pyrit.memory import CentralMemory, data_serializer_factory from pyrit.models import ( - MEDIA_PATH_DATA_TYPES, AtomicAttackIdentifier, AttackIdentifier, AttackTechniqueIdentifier, @@ -419,52 +418,8 @@ def _replace_attack_in_atomic( @staticmethod async def _persist_base64_pieces_async(request: AddMessageRequest) -> None: - """ - Resolve original and converted media independently, updating values in-place. - - The frontend sends binary media (images, audio, etc.) as base64 strings - with a ``*_path`` data_type. The PyRIT target layer expects ``*_path`` - values to be **file paths**, so we decode the base64 data, write it to - the results store, and replace the request values with the resulting - file path before the message is built. - - If the value is already an HTTP(S) URL (e.g. an Azure Blob Storage URL - from a remixed/copied message), it is kept as-is since the file already - exists in storage. - """ - for piece in request.pieces: - original_value = piece.original_value - converted_value = piece.converted_value - converted_type = piece.converted_value_data_type or piece.data_type - if piece.data_type in MEDIA_PATH_DATA_TYPES: - result = await persist_media_value_async( - value=original_value, - data_type=piece.data_type, - mime_type=piece.mime_type, - serializer_factory=data_serializer_factory, - ) - if result.resolved: - original_value = result.value - if converted_value is None or ( - converted_value == piece.original_value and converted_type == piece.data_type - ): - converted_value = original_value - - if ( - converted_value is not None - and converted_type in MEDIA_PATH_DATA_TYPES - and (converted_value != original_value or converted_type != piece.data_type) - ): - result = await persist_media_value_async( - value=converted_value, - data_type=converted_type, - serializer_factory=data_serializer_factory, - ) - if result.resolved: - converted_value = result.value - - piece.original_value = original_value - piece.converted_value = converted_value + """Persist original and converted media before sending or storing a message.""" + await persist_message_pieces_async(pieces=request.pieces, serializer_factory=data_serializer_factory) async def _send_and_store_message_async( self, diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index d57188f2c7..4e6a9f8c5e 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -152,11 +152,12 @@ async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. - Returns every constructible target with its derived constructor + Returns every registered target with its derived constructor parameters and the auth modes it supports, all projected from the - registry's ``TargetMetadata``. Deciding which entries to surface to a - user is a presentation concern owned by the caller (e.g. the frontend), - not this service. + registry's ``TargetMetadata``. Types that read local files or load model + code are listed but cannot be created through the API. Deciding which + entries to surface to a user is a presentation concern owned by the + caller (e.g. the frontend), not this service. Returns: TargetTypeResponse containing all available target classes. @@ -181,9 +182,10 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn reference resolution, and construction are owned by the ``TargetRegistry``. Endpoint trust and identity token minting are owned by the target classes themselves. This service only enforces the - request-level auth contract: for ``identity`` it confirms the target - supports it and omits the api_key so the target validates its own - endpoint and authenticates itself. + request-level contract: it rejects target types that read local files or + load model code, and for ``identity`` it confirms the target supports it + and omits the api_key so the target validates its own endpoint and + authenticates itself. Args: request: The create target request with type, params, and auth_mode. @@ -192,10 +194,11 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn TargetInstance with the new target's details. Raises: - ValueError: If the target type is not registered or identity auth is - requested but unsupported by the target type. Construction errors - (unknown params, incompatible inner targets, unrecognized identity - endpoints) are raised by the registry / target classes. + ValueError: If the target type is not registered, uses local files or + model code, or identity auth is requested but unsupported by the + target type. Construction errors (unknown params, incompatible inner + targets, unrecognized identity endpoints) are raised by the + registry / target classes. """ if request.type not in self._registry: raise ValueError( @@ -203,6 +206,11 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn ) target_cls = self._registry.get_class(request.type) + if target_cls.uses_host_resources: + raise ValueError( + f"Target type '{request.type}' reads local files or loads model code and cannot be created " + "through the API. Register it in Python or with an initializer instead." + ) params: dict[str, Any] = dict(request.params) if request.auth_mode == "identity": diff --git a/pyrit/prompt_target/common/prompt_target.py b/pyrit/prompt_target/common/prompt_target.py index 3c58d2a34d..5aa6c076af 100644 --- a/pyrit/prompt_target/common/prompt_target.py +++ b/pyrit/prompt_target/common/prompt_target.py @@ -84,6 +84,11 @@ class PromptTarget(Identifiable): # Azure Blob Storage, Prompt Shield) override this to add ``"identity"``. supported_auth_modes: ClassVar[tuple[AuthMode, ...]] = ("api_key",) + # Declarative fact consumed by the create-target service. Targets that read files + # from the machine running PyRIT or load model code on it set this to True, and the + # create-target API rejects them; register those targets in Python or with an initializer. + uses_host_resources: ClassVar[bool] = False + def __init_subclass__(cls, **kwargs: object) -> None: """ Validate that subclasses follow the keyword-only ``__init__`` contract. diff --git a/pyrit/prompt_target/http_target/httpx_api_target.py b/pyrit/prompt_target/http_target/httpx_api_target.py index cf13f7897f..12bb912fbd 100644 --- a/pyrit/prompt_target/http_target/httpx_api_target.py +++ b/pyrit/prompt_target/http_target/httpx_api_target.py @@ -6,7 +6,7 @@ import mimetypes from collections.abc import Callable from pathlib import Path -from typing import Any, Literal +from typing import Any, ClassVar, Literal import aiofiles import httpx @@ -37,6 +37,8 @@ class HTTPXAPITarget(HTTPTarget): """ _PATH_TYPES: frozenset[str] = frozenset({"image_path", "audio_path", "video_path", "binary_path"}) + # Uploads files from the local file system. + uses_host_resources: ClassVar[bool] = True _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( capabilities=TargetCapabilities( supports_multi_turn=True, diff --git a/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py b/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py index db37e5de81..e42767f4b0 100644 --- a/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py +++ b/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py @@ -6,7 +6,7 @@ import logging import warnings from pathlib import Path -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, ClassVar, cast from pyrit.common import default_values from pyrit.exceptions import EmptyResponseException, pyrit_target_retry @@ -36,6 +36,9 @@ class HuggingFaceChatTarget(PromptTarget): ) ) + # Loads models, and optionally their code, on the local machine. + uses_host_resources: ClassVar[bool] = True + # Class-level cache for model and tokenizer _cached_model: Any = None _cached_tokenizer: Any = None diff --git a/tests/unit/backend/test_api_routes.py b/tests/unit/backend/test_api_routes.py index e512b8e6cc..fb7e0334a0 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -256,6 +256,20 @@ def test_create_attack_target_not_found(self, client: TestClient) -> None: assert response.status_code == status.HTTP_404_NOT_FOUND + def test_create_attack_invalid_media_returns_bad_request(self, client: TestClient) -> None: + """Validation errors that are not lookups return 400.""" + with patch("pyrit.backend.routes.attacks.get_attack_service") as mock_get_service: + mock_service = MagicMock() + mock_service.create_attack_async = AsyncMock( + side_effect=ValueError("Media path is outside the allowed results directory.") + ) + mock_get_service.return_value = mock_service + + response = client.post("/api/attacks", json={"target_registry_name": "target"}) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert "outside the allowed results directory" in response.json()["detail"] + def test_get_attack_success(self, client: TestClient) -> None: """Test getting an attack by ID.""" now = datetime.now(UTC) diff --git a/tests/unit/backend/test_attack_models.py b/tests/unit/backend/test_attack_models.py index 49f9d1966e..b4c89dc883 100644 --- a/tests/unit/backend/test_attack_models.py +++ b/tests/unit/backend/test_attack_models.py @@ -8,7 +8,7 @@ import pytest from pydantic import ValidationError -from pyrit.backend.models.attacks import MessagePieceRequest +from pyrit.backend.models.attacks import AddMessageRequest, MessagePieceRequest from pyrit.models import PromptDataType @@ -74,3 +74,25 @@ def test_message_piece_accepts_applied_converter_order(converter_ids: list[str]) def test_message_piece_rejects_applied_converters_without_value(converter_ids: list[str]) -> None: with pytest.raises(ValidationError, match="applied_converter_ids requires converted_value"): MessagePieceRequest(original_value="source", applied_converter_ids=converter_ids) + + +@pytest.mark.parametrize( + "piece", + [ + MessagePieceRequest(data_type="url", original_value="https://example.test/image.png"), + MessagePieceRequest( + original_value="source", converted_value="https://example.test/image.png", converted_value_data_type="url" + ), + ], +) +def test_add_message_rejects_url_pieces_when_sending(piece: MessagePieceRequest) -> None: + with pytest.raises(ValidationError, match="URL pieces cannot be sent"): + AddMessageRequest(pieces=[piece], send=True, target_conversation_id="conversation") + + +def test_add_message_stores_url_pieces_without_sending() -> None: + piece = MessagePieceRequest(data_type="url", original_value="https://example.test/blob.png") + + request = AddMessageRequest(role="assistant", pieces=[piece], send=False, target_conversation_id="conversation") + + assert request.pieces[0].data_type == "url" diff --git a/tests/unit/backend/test_attack_service.py b/tests/unit/backend/test_attack_service.py index a28a30421f..6ace7a4c33 100644 --- a/tests/unit/backend/test_attack_service.py +++ b/tests/unit/backend/test_attack_service.py @@ -12,6 +12,7 @@ import json import uuid from datetime import UTC, datetime +from pathlib import Path from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -888,6 +889,63 @@ async def test_create_attack_stores_prepended_conversation(self, attack_service, mock_memory.add_attack_results_to_memory_async.assert_called_once() mock_memory.add_message_pieces_to_memory_async.assert_called() + async def test_create_attack_rejects_prepended_media_outside_results_before_writing( + self, attack_service, mock_memory, tmp_path: Path + ) -> None: + """Prepended media must point into managed storage, and nothing is written when it does not.""" + mock_memory.results_path = str(tmp_path / "results") + outside_image = tmp_path / "outside.png" + outside_image.write_bytes(b"PNG") + prepended = [ + PrependedMessageRequest( + role="user", + pieces=[MessagePieceRequest(data_type="image_path", original_value=str(outside_image))], + ) + ] + + with ( + patch("pyrit.backend.services.attack_service.get_target_service") as mock_get_target_service, + pytest.raises(ValueError, match="outside the allowed results directory"), + ): + mock_target_service = MagicMock() + mock_target_service.get_target_async = AsyncMock(return_value=MagicMock(type="TextTarget")) + mock_get_target_service.return_value = mock_target_service + await attack_service.create_attack_async( + request=CreateAttackRequest(target_registry_name="target-1", prepended_conversation=prepended) + ) + + mock_memory.add_attack_results_to_memory_async.assert_not_called() + mock_memory.add_message_pieces_to_memory_async.assert_not_called() + + async def test_create_attack_persists_prepended_base64_media(self, attack_service, mock_memory) -> None: + """Prepended base64 media is written to result storage and stored as a file path.""" + serializer = MagicMock(value="/results/prompt-memory-entries/images/prepended.png") + serializer.save_b64_image_async = AsyncMock() + prepended = [ + PrependedMessageRequest( + role="user", + pieces=[MessagePieceRequest(data_type="image_path", original_value="aW1hZ2U=", mime_type="image/png")], + ) + ] + + with ( + patch("pyrit.backend.services.attack_service.get_target_service") as mock_get_target_service, + patch("pyrit.backend.services.attack_service.data_serializer_factory", return_value=serializer) as factory, + ): + mock_target_service = MagicMock() + mock_target_service.get_target_async = AsyncMock(return_value=MagicMock(type="TextTarget")) + mock_target_service.get_target_object.return_value.get_identifier.return_value = ComponentIdentifier( + class_name="TextTarget", class_module="pyrit.prompt_target" + ) + mock_get_target_service.return_value = mock_target_service + await attack_service.create_attack_async( + request=CreateAttackRequest(target_registry_name="target-1", prepended_conversation=prepended) + ) + + assert factory.call_args.kwargs["category"] == "prompt-memory-entries" + stored_piece = mock_memory.add_message_pieces_to_memory_async.call_args.kwargs["message_pieces"][0] + assert stored_piece.original_value == "/results/prompt-memory-entries/images/prepended.png" + async def test_create_attack_lowers_system_prompt_to_system_message(self, attack_service, mock_memory) -> None: """Test that system_prompt is lowered to a single system-role message at sequence 0.""" with patch("pyrit.backend.services.attack_service.get_target_service") as mock_get_target_service: diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index 26497ac54b..f50c18cf05 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -38,6 +38,14 @@ from pyrit.prompt_normalizer import PromptNormalizer from pyrit.registry.components import ConverterRegistry +_BLOB_RESULTS_ROOT = "https://account.blob.core.windows.net/results" +_STORED_BLOB_IMAGE = f"{_BLOB_RESULTS_ROOT}/prompt-memory-entries/images/image.png" + + +def _results_root(root: str): + return patch.object(CentralMemory, "get_memory_instance", return_value=MagicMock(results_path=root)) + + _TOKEN_BIJECTION_VOCAB = ( "cat", "dog", @@ -600,10 +608,11 @@ async def test_create_with_path_or_str_upload( 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}) - ) + url = f"{_BLOB_RESULTS_ROOT}/prompt-memory-entries/inputs/input.{extension}" + with _results_root(_BLOB_RESULTS_ROOT): + 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 @@ -612,6 +621,24 @@ async def test_create_with_path_or_str_url( assert list(upload_service._upload_path.iterdir()) == [] assert await upload_service.delete_converter_async(converter_id=response.converter_id) + @pytest.mark.parametrize( + "url", + [ + "https://account.blob.core.windows.net/container/input.mp4", + f"{_BLOB_RESULTS_ROOT}/private/input.mp4", + f"{_BLOB_RESULTS_ROOT}/prompt-memory-entries/inputs/..\\..\\private/input.mp4", + f"{_BLOB_RESULTS_ROOT}%5Cprompt-memory-entries/private/input.mp4", + ], + ) + async def test_path_or_str_url_outside_results_is_rejected( + self, upload_service: ConverterService, url: str + ) -> None: + with _results_root(_BLOB_RESULTS_ROOT), pytest.raises(ValueError, match="result storage"): + await upload_service.create_converter_async( + request=CreateConverterRequest(name="remote", type="AddImageVideoConverter", params={"video_path": url}) + ) + assert upload_service._registry.instances.get_entry("remote") is None + @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 @@ -914,31 +941,60 @@ async def test_preview_conversion_with_converter_ids(self) -> None: assert len(result.steps) == 1 assert result.steps[0].converter_id == "conv-1" - @pytest.mark.parametrize( - ("value", "resolved_value"), - [ - ("https://example.test/image.png", "https://example.test/image.png"), - ("/api/media?path=%2Ftmp%2Fimage.png", "/tmp/image.png"), - ], - ) - async def test_preview_conversion_resolves_reference_without_persistence( - self, value: str, resolved_value: str - ) -> None: - """Remote and local media references bypass serializer persistence.""" + async def test_preview_conversion_keeps_stored_blob_url_without_persistence(self) -> None: + """Blob URLs inside the results container bypass serializer persistence.""" service = ConverterService() request = ConverterPreviewRequest( - original_value=value, + original_value=_STORED_BLOB_IMAGE, original_value_data_type="image_path", converter_ids=[], ) - with patch("pyrit.backend.services.converter_service.data_serializer_factory") as factory: + with ( + _results_root(_BLOB_RESULTS_ROOT), + patch("pyrit.backend.services.converter_service.data_serializer_factory") as factory, + ): result = await service.preview_conversion_async(request=request) - assert result.original_value == value - assert result.converted_value == resolved_value + assert result.original_value == _STORED_BLOB_IMAGE + assert result.converted_value == _STORED_BLOB_IMAGE factory.assert_not_called() + async def test_preview_conversion_resolves_media_reference_without_persistence(self, tmp_path: Path) -> None: + """``/api/media`` references inside the results directory bypass serializer persistence.""" + service = ConverterService() + stored_path = (tmp_path / "prompt-memory-entries" / "image.png").resolve() + request = ConverterPreviewRequest( + original_value=f"/api/media?path={stored_path}", + original_value_data_type="image_path", + converter_ids=[], + ) + + with ( + _results_root(str(tmp_path)), + patch("pyrit.backend.services.converter_service.data_serializer_factory") as factory, + ): + result = await service.preview_conversion_async(request=request) + + assert result.converted_value == str(stored_path) + factory.assert_not_called() + + @pytest.mark.parametrize("value", ["https://example.test/image.png", "/api/media?path=/etc/hostname"]) + async def test_preview_conversion_rejects_media_outside_results(self, tmp_path: Path, value: str) -> None: + """Media values must point into this server's media storage.""" + service = ConverterService() + request = ConverterPreviewRequest(original_value=value, original_value_data_type="image_path", converter_ids=[]) + + with _results_root(str(tmp_path)), pytest.raises(ValueError, match="result storage|results directory"): + await service.preview_conversion_async(request=request) + + def test_preview_request_rejects_url_input(self) -> None: + """Preview never downloads caller-supplied URLs.""" + with pytest.raises(ValidationError, match="URL input is not supported"): + ConverterPreviewRequest( + original_value="https://example.test/image.png", original_value_data_type="url", converter_ids=[] + ) + async def test_preview_conversion_chains_multiple_converters(self) -> None: """Test that preview chains multiple converters.""" service = ConverterService() @@ -977,7 +1033,7 @@ async def test_preview_conversion_chains_multiple_converters(self) -> None: [ (" source \n", "text", "", "text", " transformed \n", "text"), (" source \n", "text", "generated.png", "image_path", "edited.png", "image_path"), - ("https://example.test/image.png", "image_path", "converted.wav", "audio_path", "caption", "text"), + (_STORED_BLOB_IMAGE, "image_path", "converted.wav", "audio_path", "caption", "text"), ], ) async def test_preview_uses_normalizer_without_sending_or_storing_async( @@ -1004,6 +1060,7 @@ async def test_preview_uses_normalizer_without_sending_or_storing_async( ) with ( + _results_root(_BLOB_RESULTS_ROOT), patch("pyrit.backend.services.converter_service.PromptNormalizer", return_value=normalizer), patch.object(normalizer, "convert_values_async", wraps=normalizer.convert_values_async) as convert, patch.object(normalizer, "send_prompt_async", new_callable=AsyncMock) as send, @@ -1105,11 +1162,14 @@ async def test_preview_conversion_unmarked_media_retains_result_type_async( instance = Base64Converter() upload_service._registry.instances.register(instance, name="media") request = ConverterPreviewRequest( - original_value="https://example.test/image.png", + original_value=_STORED_BLOB_IMAGE, original_value_data_type="image_path", converter_ids=["media"], ) - with patch.object(instance, "convert_async", new_callable=AsyncMock) as convert: + with ( + _results_root(_BLOB_RESULTS_ROOT), + patch.object(instance, "convert_async", new_callable=AsyncMock) as convert, + ): convert.return_value = converter.ConverterResult(output_text="converted.wav", output_type="audio_path") result = await upload_service.preview_conversion_async(request=request) convert.assert_awaited_once_with(prompt=request.original_value, input_type="image_path") @@ -1254,9 +1314,10 @@ async def test_preview_conversion_propagates_invalid_base64_error_after_path_fai await service.preview_conversion_async(request=request) async def test_preview_conversion_preserves_existing_file(self, tmp_path: Path) -> None: - """Existing local media paths pass through without being persisted again.""" + """Existing media files in the results directory pass through without being persisted again.""" service = ConverterService() - media_path = tmp_path / "input.wav" + media_path = tmp_path / "prompt-memory-entries" / "audio" / "input.wav" + media_path.parent.mkdir(parents=True) media_path.write_bytes(b"RIFF") request = ConverterPreviewRequest( original_value=str(media_path), @@ -1264,11 +1325,14 @@ async def test_preview_conversion_preserves_existing_file(self, tmp_path: Path) converter_ids=[], ) - with patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory: + with ( + _results_root(str(tmp_path)), + patch("pyrit.backend.services.converter_service.data_serializer_factory") as mock_factory, + ): result = await service.preview_conversion_async(request=request) mock_factory.assert_not_called() - assert result.converted_value == str(media_path) + assert result.converted_value == str(media_path.resolve()) class TestGetConverterObjectsForIds: diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index 2fdab2f1af..202c913a8d 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -5,10 +5,20 @@ from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import quote import pytest -from pyrit.backend.services.media_persistence import MediaOrigin, persist_media_value_async +from pyrit.backend.services.media_persistence import ( + MEDIA_SUBDIRECTORIES, + MediaOrigin, + persist_media_value_async, + require_managed_blob_url, +) +from pyrit.memory import CentralMemory +from pyrit.memory.storage.storage import AzureBlobStorageIO + +_BLOB_ROOT = "https://account.blob.core.windows.net/results" def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: @@ -18,41 +28,193 @@ def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: return serializer +def _results_root(root: str): + return patch.object(CentralMemory, "get_memory_instance", return_value=MagicMock(results_path=root)) + + +@pytest.fixture +def stored_image(tmp_path: Path) -> Path: + image_path = tmp_path / "results" / "prompt-memory-entries" / "images" / "stored.png" + image_path.parent.mkdir(parents=True) + image_path.write_bytes(b"PNG") + return image_path + + +async def test_media_reference_inside_results_is_resolved(stored_image: Path) -> None: + factory = MagicMock() + root = stored_image.parents[2] + + with _results_root(str(root)): + result = await persist_media_value_async( + value=f"/api/media?path={quote(str(stored_image))}", data_type="image_path", serializer_factory=factory + ) + + assert result.origin is MediaOrigin.MEDIA_REFERENCE + assert result.value == str(stored_image.resolve()) + assert result.resolved is True + assert result.persisted is False + factory.assert_not_called() + + +async def test_existing_local_path_inside_results_is_not_persisted(stored_image: Path) -> None: + factory = MagicMock() + + with _results_root(str(stored_image.parents[2])): + result = await persist_media_value_async( + value=str(stored_image), data_type="image_path", serializer_factory=factory + ) + + assert result.origin is MediaOrigin.LOCAL_PATH + assert result.value == str(stored_image.resolve()) + assert result.persisted is False + factory.assert_not_called() + + +async def test_blob_url_inside_results_container_is_kept() -> None: + value = f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png?sv=2024&sig=signature" + + with _results_root(_BLOB_ROOT): + result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) + + assert result.origin is MediaOrigin.REMOTE_URL + assert result.value == value + assert result.persisted is False + + @pytest.mark.parametrize( - ("value", "origin", "resolved_value", "resolved"), + "value", [ - ("https://example.test/media.png", MediaOrigin.REMOTE_URL, "https://example.test/media.png", True), - ("/api/media?path=%2Ftmp%2Fmedia.png", MediaOrigin.MEDIA_REFERENCE, "/tmp/media.png", True), - ("/api/media", MediaOrigin.MEDIA_REFERENCE, "/api/media", False), + "https://localHoSt/", + "https://127.0.1.2/", + "https://0177.0.23.19/", + "https://2130706433/", + "https://0x7f.00331.0246.174/", + "https://[::1]/", + "https://[fc00::]/", + "https://169.254.169.254/", + "https://example.test/media.png", + "https://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", ], ) -async def test_existing_references_are_not_persisted( - value: str, origin: MediaOrigin, resolved_value: str, resolved: bool -) -> None: +async def test_url_is_rejected_when_results_are_local(stored_image: Path, value: str) -> None: factory = MagicMock() - result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + with _results_root(str(stored_image.parents[2])), pytest.raises(ValueError, match="result storage"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) - assert result.origin is origin - assert result.value == resolved_value - assert result.resolved is resolved - assert result.persisted is False factory.assert_not_called() -async def test_existing_local_path_is_not_persisted(tmp_path: Path) -> None: - media_path = tmp_path / "audio.wav" - media_path.write_bytes(b"RIFF") +@pytest.mark.parametrize( + "value", + [ + "https://other.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}-other/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}/other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/../other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/images/..\\..\\other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/images/%5C..%5C..%5Cother-folder/stored.png", + "http://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", + "https://169.254.169.254/results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}%2Fprompt-memory-entries/other-folder/stored.png", + f"{_BLOB_ROOT}%2fprompt-memory-entries/stored.png", + f"{_BLOB_ROOT}%5Cprompt-memory-entries/stored.png", + f"{_BLOB_ROOT}/./prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}//prompt-memory-entries/images/stored.png", + "https://account.blob.core.windows.net//results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries", + ], +) +async def test_url_outside_results_container_is_rejected(value: str) -> None: + with _results_root(_BLOB_ROOT), pytest.raises(ValueError, match="result storage"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) + + +@pytest.mark.parametrize( + "value", + [ + f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}/seed-prompt-entries/images/stored.png?sv=2024&sig=signature", + f"{_BLOB_ROOT}/prompt-memory-entries\\images\\stored.png", + ], +) +def test_accepted_blob_url_is_read_from_a_media_folder(value: str) -> None: + with _results_root(_BLOB_ROOT): + require_managed_blob_url(value) + + blob_segments = AzureBlobStorageIO(container_url=_BLOB_ROOT)._resolve_blob_name(value).split("/") + assert blob_segments[0] in MEDIA_SUBDIRECTORIES + assert ".." not in blob_segments + + +def test_blob_url_generated_under_trailing_slash_root_is_accepted() -> None: + root = f"{_BLOB_ROOT}/" + + with _results_root(root): + require_managed_blob_url(f"{root}/prompt-memory-entries/images/stored.png") + for value in (f"{_BLOB_ROOT}/other-folder/stored.png", f"{root}/other-folder/stored.png"): + with pytest.raises(ValueError, match="result storage"): + require_managed_blob_url(value) + + +@pytest.mark.parametrize("root", ["https://account.blob.core.windows.net", "https://account.blob.core.windows.net/"]) +def test_blob_url_is_rejected_when_results_root_has_no_container(root: str) -> None: + with _results_root(root), pytest.raises(ValueError, match="result storage"): + require_managed_blob_url("https://account.blob.core.windows.net/prompt-memory-entries/images/stored.png") + + +@pytest.mark.parametrize("path_suffix", ["/../../OtherPath/", "/..//OtherPath/", "/%2E%2E%2f/OtherPath/"]) +async def test_media_reference_traversal_is_rejected(stored_image: Path, path_suffix: str) -> None: + root = stored_image.parents[2] + value = f"/api/media?path={quote(str(root / 'prompt-memory-entries'))}{path_suffix}" + + with _results_root(str(root)), pytest.raises(ValueError, match="results directory"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) + + +async def test_local_path_outside_results_is_rejected(stored_image: Path, tmp_path: Path) -> None: + outside_file = tmp_path / "outside.txt" + outside_file.write_text("outside") factory = MagicMock() - result = await persist_media_value_async(value=str(media_path), data_type="audio_path", serializer_factory=factory) + with _results_root(str(stored_image.parents[2])), pytest.raises(ValueError, match="outside the allowed results"): + await persist_media_value_async(value=str(outside_file), data_type="binary_path", serializer_factory=factory) - assert result.origin is MediaOrigin.LOCAL_PATH - assert result.value == str(media_path) - assert result.persisted is False factory.assert_not_called() +async def test_local_path_outside_media_folders_is_rejected(stored_image: Path) -> None: + root = stored_image.parents[2] + database_file = root / "pyrit.db" + database_file.write_bytes(b"db") + + with _results_root(str(root)), pytest.raises(ValueError, match="not in a media folder"): + await persist_media_value_async( + value=str(database_file), data_type="binary_path", serializer_factory=MagicMock() + ) + + +async def test_symlink_escaping_results_is_rejected(stored_image: Path, tmp_path: Path) -> None: + outside_file = tmp_path / "outside.png" + outside_file.write_bytes(b"PNG") + link = stored_image.parent / "link.png" + link.symlink_to(outside_file) + + with _results_root(str(stored_image.parents[2])), pytest.raises(ValueError, match="outside the allowed results"): + await persist_media_value_async(value=str(link), data_type="image_path", serializer_factory=MagicMock()) + + +async def test_local_path_is_rejected_when_results_are_remote(stored_image: Path) -> None: + with _results_root(_BLOB_ROOT), pytest.raises(ValueError, match="Azure Blob Storage"): + await persist_media_value_async(value=str(stored_image), data_type="image_path", serializer_factory=MagicMock()) + + +@pytest.mark.parametrize("value", ["/api/media", "/api/media?path=", "/api/media?path=a.png&path=b.png"]) +async def test_malformed_media_reference_is_rejected(value: str) -> None: + with pytest.raises(ValueError, match="exactly one path"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) + + async def test_data_uri_uses_explicit_mime_before_uri_mime() -> None: serializer = _serializer(value="/saved/media.jpg") factory = MagicMock(return_value=serializer) diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 6a436df453..e4894b2d94 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -57,6 +57,15 @@ def mock_memory(patch_central_database: MagicMock) -> Iterator[MagicMock]: yield memory +def _stored_media(value: str, *, root: Path) -> str: + """Expand ``media:`` to an ``/api/media`` reference and ``stored:`` to its stored path.""" + kind, _, name = value.partition(":") + stored_path = root / "prompt-memory-entries" / name + if kind == "media": + return f"/api/media?path={stored_path}" + return str(stored_path) if kind == "stored" else value + + @pytest.fixture def message_send_service(mock_memory: MagicMock) -> MessageSendService: return MessageSendService(scheduler=ManualSendScheduler()) @@ -444,6 +453,7 @@ async def test_add_message_preserves_converter_configuration_targeting( self, message_send_service, mock_memory ) -> None: """Test that request and response converter targeting reaches the normalizer.""" + mock_memory.results_path = "https://account.blob.core.windows.net/results" ar = make_attack_result(conversation_id="test-id") mock_memory.get_attack_results_async.return_value = [ar] mock_memory.get_message_pieces_async.return_value = [] @@ -502,7 +512,7 @@ async def test_add_message_preserves_converter_configuration_targeting( MessagePieceRequest(original_value="Hello"), MessagePieceRequest( data_type="image_path", - original_value="https://example.com/image.png", + original_value="https://account.blob.core.windows.net/results/prompt-memory-entries/image.png", ), ], target_conversation_id="test-id", @@ -778,17 +788,17 @@ async def test_persists_original_and_converted_media_independently_async( for serializer in serializers: serializer.save_b64_image_async.assert_awaited_once() - @pytest.mark.parametrize( - ("converted_value", "expected_value"), - [ - ("/api/media?path=preview.png", "preview.png"), - ("https://example.com/preview.png?token=example", "https://example.com/preview.png?token=example"), - ("preview.png", "preview.png"), - ], - ) + @pytest.mark.parametrize("stored_in_blob", [False, True]) async def test_converted_media_references_are_not_repersisted_async( - self, *, converted_value: str, expected_value: str + self, *, mock_memory: MagicMock, tmp_path: Path, stored_in_blob: bool ) -> None: + stored_path = (tmp_path / "prompt-memory-entries" / "preview.png").resolve() + if stored_in_blob: + mock_memory.results_path = "https://account.blob.core.windows.net/results" + converted_value = expected_value = f"{mock_memory.results_path}/prompt-memory-entries/preview.png?sv=1" + else: + mock_memory.results_path = str(tmp_path) + converted_value, expected_value = f"/api/media?path={stored_path}", str(stored_path) request = AddMessageRequest( pieces=[ MessagePieceRequest( @@ -800,16 +810,33 @@ async def test_converted_media_references_are_not_repersisted_async( send=False, target_conversation_id="test-id", ) - with ( - patch("pyrit.backend.services.media_persistence.Path.is_file", return_value=True), - patch("pyrit.backend.services.message_send_service.data_serializer_factory") as factory, - ): + with patch("pyrit.backend.services.message_send_service.data_serializer_factory") as factory: await MessageSendService._persist_base64_pieces_async(request) assert request.pieces[0].original_value == "source" assert request.pieces[0].converted_value == expected_value factory.assert_not_called() + @pytest.mark.parametrize("converted_value", ["https://example.com/preview.png", "/api/media?path=/etc/hostname"]) + async def test_converted_media_outside_results_is_rejected_async( + self, *, mock_memory: MagicMock, tmp_path: Path, converted_value: str + ) -> None: + mock_memory.results_path = str(tmp_path) + request = AddMessageRequest( + pieces=[ + MessagePieceRequest( + original_value="source", + converted_value=converted_value, + converted_value_data_type="image_path", + ) + ], + send=False, + target_conversation_id="test-id", + ) + + with pytest.raises(ValueError, match="result storage|results directory"): + await MessageSendService._persist_base64_pieces_async(request) + async def test_identical_original_and_converted_media_saved_once_async(self) -> None: request = AddMessageRequest( pieces=[ @@ -865,11 +892,14 @@ async def test_text_pieces_are_unchanged(self, message_send_service) -> None: await MessageSendService._persist_base64_pieces_async(request) assert request.pieces[0].original_value == "hello" - @pytest.mark.parametrize("converted_value", [None, "https://example.com/converted.png"]) + @pytest.mark.parametrize( + "converted_value", [None, "https://account.blob.core.windows.net/results/prompt-memory-entries/converted.png"] + ) async def test_image_piece_is_saved_to_file( - self, *, message_send_service: MessageSendService, converted_value: str | None + self, *, message_send_service: MessageSendService, mock_memory: MagicMock, converted_value: str | None ) -> None: """Base64 image data should be saved to disk and value replaced with file path.""" + mock_memory.results_path = "https://account.blob.core.windows.net/results" request = AddMessageRequest( role="user", pieces=[ @@ -1054,14 +1084,16 @@ async def test_path_data_type_supplies_extension_when_mime_type_missing(self, me ) assert request.pieces[0].original_value == "/saved/image.png" - async def test_http_url_is_kept_as_is(self, message_send_service) -> None: - """HTTPS blob URLs should not be re-persisted.""" + async def test_http_url_is_kept_as_is(self, message_send_service, mock_memory: MagicMock) -> None: + """HTTPS blob URLs inside the results container should not be re-persisted.""" + mock_memory.results_path = "https://myblob.blob.core.windows.net/results" + blob_url = "https://myblob.blob.core.windows.net/results/prompt-memory-entries/images/photo.png?sv=2024" request = AddMessageRequest( role="user", pieces=[ MessagePieceRequest( data_type="image_path", - original_value="https://myblob.blob.core.windows.net/images/photo.png?sv=2024", + original_value=blob_url, mime_type="image/png", ), ], @@ -1071,17 +1103,21 @@ async def test_http_url_is_kept_as_is(self, message_send_service) -> None: await MessageSendService._persist_base64_pieces_async(request) - assert request.pieces[0].original_value == ("https://myblob.blob.core.windows.net/images/photo.png?sv=2024") + assert request.pieces[0].original_value == blob_url assert request.pieces[0].converted_value == request.pieces[0].original_value - async def test_media_reference_is_resolved_without_persistence(self, message_send_service) -> None: - """Local media URLs are converted back to their decoded file paths.""" + async def test_media_reference_is_resolved_without_persistence( + self, message_send_service, mock_memory: MagicMock, tmp_path: Path + ) -> None: + """Local media URLs are converted back to their canonical file paths.""" + mock_memory.results_path = str(tmp_path) + stored_path = (tmp_path / "prompt-memory-entries" / "image.png").resolve() request = AddMessageRequest( role="user", pieces=[ MessagePieceRequest( data_type="image_path", - original_value="/api/media?path=%2Ftmp%2Fimage.png", + original_value=f"/api/media?path={stored_path}", ), ], send=False, @@ -1091,13 +1127,17 @@ async def test_media_reference_is_resolved_without_persistence(self, message_sen with patch("pyrit.backend.services.message_send_service.data_serializer_factory") as factory: await MessageSendService._persist_base64_pieces_async(request) - assert request.pieces[0].original_value == "/tmp/image.png" - assert request.pieces[0].converted_value == "/tmp/image.png" + assert request.pieces[0].original_value == str(stored_path) + assert request.pieces[0].converted_value == str(stored_path) factory.assert_not_called() - async def test_existing_file_is_kept_without_persistence(self, message_send_service, tmp_path: Path) -> None: - """An existing path remains the canonical original and converted value.""" - media_path = tmp_path / "image.png" + async def test_existing_file_is_kept_without_persistence( + self, message_send_service, mock_memory: MagicMock, tmp_path: Path + ) -> None: + """An existing stored file remains the canonical original and converted value.""" + mock_memory.results_path = str(tmp_path) + media_path = tmp_path / "prompt-memory-entries" / "images" / "image.png" + media_path.parent.mkdir(parents=True) media_path.write_bytes(b"image") request = AddMessageRequest( role="user", @@ -1109,10 +1149,27 @@ async def test_existing_file_is_kept_without_persistence(self, message_send_serv with patch("pyrit.backend.services.message_send_service.data_serializer_factory") as factory: await MessageSendService._persist_base64_pieces_async(request) - assert request.pieces[0].original_value == str(media_path) - assert request.pieces[0].converted_value == str(media_path) + assert request.pieces[0].original_value == str(media_path.resolve()) + assert request.pieces[0].converted_value == str(media_path.resolve()) factory.assert_not_called() + async def test_existing_file_outside_results_is_rejected( + self, message_send_service, mock_memory: MagicMock, tmp_path: Path + ) -> None: + """Files outside the results media folders cannot be used as message media.""" + mock_memory.results_path = str(tmp_path / "results") + outside_path = tmp_path / "outside.png" + outside_path.write_bytes(b"image") + request = AddMessageRequest( + role="user", + pieces=[MessagePieceRequest(data_type="image_path", original_value=str(outside_path))], + send=False, + target_conversation_id="test-id", + ) + + with pytest.raises(ValueError, match="outside the allowed results directory"): + await MessageSendService._persist_base64_pieces_async(request) + async def test_non_path_data_types_are_skipped(self, message_send_service) -> None: """Non *_path types like reasoning, url, function_call should not be decoded.""" request = AddMessageRequest( @@ -2428,22 +2485,22 @@ async def test_exact_preview_provenance_survives_memory_and_response_mapping_asy [ ("text", "source", "text", "", "source", ""), ("text", "", "text", "Edited preview", "", "Edited preview"), - ("text", "source", "image_path", "/api/media?path=preview.png", "source", "preview.png"), + ("text", "source", "image_path", "media:preview.png", "source", "stored:preview.png"), ( "image_path", - "/api/media?path=source.png", + "media:source.png", "text", "Exact description", - "source.png", + "stored:source.png", "Exact description", ), ( "image_path", - "/api/media?path=source.png", + "media:source.png", "audio_path", - "/api/media?path=preview.wav", - "source.png", - "preview.wav", + "media:preview.wav", + "stored:source.png", + "stored:preview.wav", ), ], ) @@ -2452,6 +2509,7 @@ async def test_send_preserves_exact_preview_and_converts_other_piece_async( *, message_send_service: MessageSendService, mock_memory: MagicMock, + tmp_path: Path, original_type: PromptDataType, original_value: str, converted_type: PromptDataType, @@ -2459,6 +2517,16 @@ async def test_send_preserves_exact_preview_and_converts_other_piece_async( expected_original: str, expected_final: str, ) -> None: + root = tmp_path.resolve() + mock_memory.results_path = str(root) + original_value, converted_value = ( + _stored_media(original_value, root=root), + _stored_media(converted_value, root=root), + ) + expected_original, expected_final = ( + _stored_media(expected_original, root=root), + _stored_media(expected_final, root=root), + ) mock_memory.get_attack_results_async.return_value = [make_attack_result(conversation_id="test-id")] preview_converter = MagicMock(spec=Converter) preview_converter.get_identifier.return_value = ComponentIdentifier( diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index ac997198b2..c7e5e0c226 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -373,6 +373,27 @@ async def test_create_target_raises_for_invalid_type(self) -> None: with pytest.raises(ValueError, match="not found"): await service.create_target_async(request=request) + @pytest.mark.parametrize( + ("target_type", "params"), + [ + ("HTTPXAPITarget", {"http_url": "http://localhost:8080/upload"}), + ("HuggingFaceChatTarget", {"model_id": "example/model"}), + ], + ) + async def test_create_target_rejects_types_using_host_resources( + self, sqlite_instance, target_type: str, params: dict[str, str] + ) -> None: + service = TargetService() + request = CreateTargetRequest(name="host-target", type=target_type, params=params) + + with ( + patch.object(service._registry, "create_named_instance") as create, + pytest.raises(ValueError, match="cannot be created through the API"), + ): + await service.create_target_async(request=request) + + create.assert_not_called() + async def test_create_target_success(self, sqlite_instance) -> None: """Test successful target creation.""" service = TargetService() From 530ddfb7b34dd40307a71d90e8ff5c2c302c1ccc Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Fri, 2 Oct 2026 20:13:01 +0000 Subject: [PATCH 2/7] Import media URLs into managed storage instead of rejecting them Media values, url pieces, and converter file parameters that are http(s) URLs are now downloaded once into prompt-memory-entries with a time limit, a size limit, and a redirect limit, and only the stored copy reaches converters and targets. Redirects are followed one hop at a time, so every hop must be a plain http(s) URL and redirect bodies are never read. url pieces take the media type of their content. Blob URLs inside the results container stay as references, without their query string. allow_media_url_import in .pyrit_conf turns URL import off. --- .pyrit_conf_example | 8 + pyrit/backend/README.md | 24 +- pyrit/backend/models/attacks.py | 18 -- pyrit/backend/models/converters.py | 18 +- pyrit/backend/services/converter_service.py | 37 ++-- pyrit/backend/services/media_persistence.py | 149 +++++++++++-- pyrit/backend/services/media_url_import.py | 172 +++++++++++++++ pyrit/backend/services/runtime_lifecycle.py | 2 + pyrit/setup/configuration_loader.py | 5 + tests/unit/backend/test_attack_models.py | 24 +- tests/unit/backend/test_converter_service.py | 96 ++++++-- tests/unit/backend/test_media_persistence.py | 174 +++++++++++++-- tests/unit/backend/test_media_url_import.py | 206 ++++++++++++++++++ .../unit/backend/test_message_send_service.py | 43 +++- tests/unit/backend/test_runtime_lifecycle.py | 13 ++ tests/unit/setup/test_configuration_loader.py | 9 + 16 files changed, 850 insertions(+), 148 deletions(-) create mode 100644 pyrit/backend/services/media_url_import.py create mode 100644 tests/unit/backend/test_media_url_import.py diff --git a/.pyrit_conf_example b/.pyrit_conf_example index 0e1cdab013..b797d8fb25 100644 --- a/.pyrit_conf_example +++ b/.pyrit_conf_example @@ -134,6 +134,14 @@ enable_live_reinitialization: false # Default: false allow_custom_initializers: false +# When true, the backend downloads http(s) media URLs from API requests (media +# values, `url` pieces, and converter file parameters) once into managed storage, +# with time, size, and redirect limits, and passes only the stored copy on. When +# false, such URLs are rejected and media must be uploaded. +# +# Default: true +allow_media_url_import: true + # Optional storage for custom initializer Python scripts. This may be a local # directory or an Azure Blob container URI with an optional blob prefix. # Container URIs may include a SAS; otherwise DefaultAzureCredential is used. Defaults to diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index 456f58db70..4f4468ee9b 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -133,13 +133,18 @@ The backend is the part of PyRIT that accepts requests from other machines, so i request values before using them: - Media values in messages, previews, and prepended conversations, and file parameters of - converters, must be uploaded content or point into this server's media storage: the - `prompt-memory-entries` and `seed-prompt-entries` folders under the memory results path. - Blob URLs are accepted only when results are stored in Azure Blob Storage, and only - inside those folders of the configured results container. Other file paths and URLs - are rejected. -- The backend only reads media URLs inside this server's result storage. Converter - previews and sent messages reject `url` pieces; stored history may still contain them. + converters, must be uploaded content, a media URL, or a reference into this server's media + storage: the `prompt-memory-entries` and `seed-prompt-entries` folders under the memory + results path. Other file paths are rejected. +- Media URLs and `url` pieces are downloaded once into that storage (60 second limit, + 100 MiB limit, at most 3 redirects, no request credentials forwarded), and only the stored + copy reaches converters and targets, so signed URLs are not passed to model providers. + This includes messages that are stored without being sent, because stored history is + replayed to targets later. A `url` piece takes the media type of its content + (`image_path`, `audio_path`, `video_path`, or `binary_path`); send a literal URL as `text`. + Blob URLs inside the configured results container are kept as references, without their + query string, instead of being downloaded. Set `allow_media_url_import: false` in + `.pyrit_conf` to reject media URLs instead. - Target types that read local files or load model code (`HTTPXAPITarget`, `HuggingFaceChatTarget`) cannot be created through the API. Register them in Python or with an initializer, where the operator controls their settings; for example, set @@ -147,8 +152,9 @@ request values before using them: Intentional exceptions: -- Target endpoints, raw HTTP requests, and their redirects are chosen by the operator and - are not restricted. Limit outbound network access in the deployment instead. +- Target endpoints, raw HTTP requests, media URLs, and their redirects are chosen by the + operator and are not restricted to particular hosts. Limit outbound network access in the + deployment instead. - Prompt content is not filtered. It is adversarial test data by design. - Any file type can be stored as a payload. `GET /api/media` only renders known image, audio, and video types inline; everything else downloads as a file. diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index 1419354ba6..0cd5a99869 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -667,24 +667,6 @@ def _validate_converter_configurations(self) -> "AddMessageRequest": return self - @model_validator(mode="after") - def _reject_url_pieces_when_sending(self) -> "AddMessageRequest": - """ - Reject URL pieces in a message that will be sent. - - Sending can route a piece through converters that download URL input on the - server. Stored history may still hold URL pieces because storing never fetches them. - - Returns: - AddMessageRequest: The validated request. - - Raises: - ValueError: If a piece to be sent has a ``url`` data type. - """ - if self.send and any("url" in (piece.data_type, piece.converted_value_data_type) for piece in self.pieces): - raise ValueError("URL pieces cannot be sent; upload the file content instead") - return self - class AddMessageResponse(BaseModel): """ diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 66a37c5ec6..316cedb42b 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -9,7 +9,7 @@ from typing import Any -from pydantic import BaseModel, Field, field_validator +from pydantic import BaseModel, Field from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN from pyrit.models import ConverterIdentifier, Parameter, PromptDataType @@ -119,22 +119,6 @@ class ConverterPreviewRequest(BaseModel): original_value_data_type: PromptDataType = Field(default="text", description="Data type of original value") converter_ids: list[str] = Field(..., description="Converter instance IDs to apply") - @field_validator("original_value_data_type") - @classmethod - def _reject_url_input(cls, value: PromptDataType) -> PromptDataType: - """ - Reject URL input so a preview never makes the server download a caller-supplied URL. - - Returns: - PromptDataType: The validated data type. - - Raises: - ValueError: If the data type is ``url``. - """ - if value == "url": - raise ValueError("URL input is not supported; upload the file content instead") - return value - class ConverterPreviewResponse(BaseModel): """Response from converter preview.""" diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index 988970efb0..f7f6c6fc74 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -37,8 +37,9 @@ CreateConverterRequest, PreviewStep, ) -from pyrit.backend.services.media_persistence import persist_media_value_async, require_managed_blob_url -from pyrit.common.azure_storage import is_azure_blob_uri +from pyrit.backend.services.media_persistence import is_managed_blob_url, persist_media_value_async +from pyrit.backend.services.media_url_import import download_media_url_async, media_extension +from pyrit.common.azure_storage import redact_url_credentials from pyrit.memory import data_serializer_factory from pyrit.models import MessagePiece, PromptDataType from pyrit.prompt_normalizer import ConverterConfiguration, PromptNormalizer @@ -223,8 +224,8 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> original_value = request.original_value data_type = request.original_value_data_type - # For path-based data types, resolve references or persist base64/data URIs. - if str(data_type).endswith("_path"): + # For path-based data types and URLs, resolve references, import URLs, or persist base64/data URIs. + if str(data_type).endswith("_path") or data_type == "url": result = await persist_media_value_async( value=original_value, data_type=data_type, @@ -236,6 +237,7 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> serializer_factory=data_serializer_factory, ) original_value = result.value + data_type = result.data_type or data_type converters = self._gather_converters(converter_ids=request.converter_ids) steps, final_value, final_type = await self._apply_converters_async( @@ -284,9 +286,9 @@ async def _persist_data_uri_params_async( 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 inside this server's - result storage, which pass through unchanged. Their data-URI uploads use the - same local storage. + An http(s) URL is downloaded once into the same local storage. ``Path | str`` + parameters also accept Azure Blob URLs inside this server's result storage, + which pass through unchanged. Inputs remain local until converter deletion or backend shutdown, even with Azure-backed memory. Converter outputs still use the configured result storage. @@ -304,8 +306,8 @@ async def _persist_data_uri_params_async( set of request-created files owned by the future registry entry. Raises: - ValueError: If a ``Path`` value is not a valid data URI, or a blob URL - points outside this server's result storage. + ValueError: If a ``Path`` value is not a data URI or an http(s) URL, or a URL + cannot be downloaded. """ metadata = self._registry.get_registered_class_metadata(converter_type) path_params = ( @@ -327,14 +329,19 @@ async def _persist_data_uri_params_async( 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): - require_managed_blob_url(value) + if isinstance(value, str) and value.startswith(("http://", "https://")): + if parameter.is_path_or_str and is_managed_blob_url(value): + result[name] = redact_url_credentials(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}") + download = await download_media_url_async(url=value) + content, extension = download.content, media_extension(download, default="") + elif isinstance(value, str) and value.startswith("data:"): + content, extension = self._decode_data_uri(parameter_name=name, data_uri=value) + else: + raise ValueError( + f"Path parameter '{name}' must be uploaded as a data URI or given as an http(s) URL" + ) - 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) diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index 2cbb024d16..30205c181b 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -17,7 +17,13 @@ from urllib.parse import parse_qs, unquote, urlparse from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS -from pyrit.common.azure_storage import is_azure_blob_uri +from pyrit.backend.services.media_url_import import ( + download_media_url_async, + media_content_type, + media_extension, + redact_url, +) +from pyrit.common.azure_storage import is_azure_blob_uri, redact_url_credentials from pyrit.memory import CentralMemory, data_serializer_factory from pyrit.models import MEDIA_PATH_DATA_TYPES @@ -27,6 +33,11 @@ # Media is only read from these folders under the memory results path. MEDIA_SUBDIRECTORIES = frozenset({"prompt-memory-entries", "seed-prompt-entries"}) +# Media type families of the path data types; binary_path accepts any content. +_MEDIA_FAMILIES: dict[PromptDataType, str] = {"image_path": "image", "audio_path": "audio", "video_path": "video"} +_CHECKED_FAMILIES = frozenset({"image", "audio", "video", "text"}) +# Data types whose request values are stored as managed media before use. +_IMPORTED_DATA_TYPES = frozenset({*MEDIA_PATH_DATA_TYPES, "url"}) class MediaOrigin(str, Enum): @@ -49,6 +60,7 @@ class MediaPersistenceResult: resolved: bool mime_type: str | None = None extension: str | None = None + data_type: PromptDataType | None = None SerializerFactory = Callable[..., Any] @@ -163,6 +175,20 @@ def require_managed_blob_url(value: str) -> None: raise ValueError("Media URLs must point to this server's result storage.") +def is_managed_blob_url(value: str) -> bool: + """ + Return whether a URL points into a media folder of the configured results container. + + Returns: + bool: True when the server can read the URL from its own result storage. + """ + try: + require_managed_blob_url(value) + except ValueError: + return False + return True + + def _media_reference_path(value: str) -> str | None: """ Return the file path from an ``/api/media?path=...`` reference. @@ -201,6 +227,63 @@ def _resolve_extension( return extension or DEFAULT_MEDIA_EXTENSIONS.get(str(data_type), ".bin") +def _downloaded_data_type(*, content_type: str | None) -> PromptDataType: + """ + Return the path data type for downloaded ``url`` content. + + Returns: + PromptDataType: ``image_path``, ``audio_path``, or ``video_path`` by media family, else ``binary_path``. + """ + family = (content_type or "").split("/", 1)[0] + for data_type, data_family in _MEDIA_FAMILIES.items(): + if family == data_family: + return data_type + return "binary_path" + + +async def _import_media_url_async( + *, + url: str, + data_type: PromptDataType, + serializer_factory: SerializerFactory, +) -> MediaPersistenceResult: + """ + Download a media URL once and store the bytes in managed media storage. + + A ``url`` value is stored under the path data type of its content. Other path types keep + their declared type and must not receive content of a different media family. + + Returns: + MediaPersistenceResult: The managed reference to the stored copy. + + Raises: + ValueError: If the download fails or the content does not match the declared media type. + """ + download = await download_media_url_async(url=url) + content_type = media_content_type(download) + resolved_type = _downloaded_data_type(content_type=content_type) if data_type == "url" else data_type + expected_family = _MEDIA_FAMILIES.get(resolved_type) + family = (content_type or "").split("/", 1)[0] + if expected_family and family in _CHECKED_FAMILIES and family != expected_family: + raise ValueError(f"Media URL {redact_url(url)} returned {content_type}, not {expected_family} content.") + extension = media_extension(download, default=DEFAULT_MEDIA_EXTENSIONS.get(str(resolved_type), ".bin")) + serializer = serializer_factory( + category="prompt-memory-entries", + data_type=resolved_type, + extension=extension, + ) + await serializer.save_data_async(download.content) + return MediaPersistenceResult( + value=str(serializer.value), + origin=MediaOrigin.REMOTE_URL, + persisted=True, + resolved=True, + mime_type=content_type, + extension=extension, + data_type=resolved_type, + ) + + async def persist_media_value_async( *, value: str, @@ -211,29 +294,39 @@ async def persist_media_value_async( serializer_factory: SerializerFactory = data_serializer_factory, ) -> MediaPersistenceResult: """ - Classify and, when needed, persist one path-typed media value. + Classify and, when needed, persist one path-typed or ``url`` media value. The two policy flags preserve the small historical differences between attack ingestion and converter preview while keeping origin detection, - extension resolution, and persistence in one component. Existing files, - ``/api/media`` references, and URLs are accepted only when they point into - this server's media storage, so callers cannot make the server read other files. + extension resolution, and persistence in one component. Existing files and + ``/api/media`` references are accepted only when they point into this + server's media storage, so callers cannot make the server read other files. + Blob URLs inside this server's result storage are kept as references, without + their query string, since the server reads them with its own credentials; any + other http(s) URL is downloaded once into managed storage, so converters and + targets only see the stored copy. Returns: A typed result containing the resolved value and persistence metadata. Raises: - ValueError: If the value names a file or URL outside the server's media storage. + ValueError: If the value names a file outside the server's media storage, or a + URL cannot be imported. """ if value.startswith(("http://", "https://")): - require_managed_blob_url(value) - return MediaPersistenceResult( - value=value, - origin=MediaOrigin.REMOTE_URL, - persisted=False, - resolved=True, - mime_type=mime_type, - ) + if is_managed_blob_url(value): + blob_type, _ = mimetypes.guess_type(urlparse(value).path, strict=False) + return MediaPersistenceResult( + value=redact_url_credentials(value), + origin=MediaOrigin.REMOTE_URL, + persisted=False, + resolved=True, + mime_type=mime_type, + data_type=_downloaded_data_type(content_type=blob_type) if data_type == "url" else data_type, + ) + return await _import_media_url_async(url=value, data_type=data_type, serializer_factory=serializer_factory) + if data_type == "url": + raise ValueError("URL pieces must use an http or https URL.") reference_path = _media_reference_path(value) if reference_path is not None: @@ -304,7 +397,9 @@ async def persist_message_pieces_async( values to be **file paths**, so base64 data is written to the results store and the request values are replaced with the resulting file path. Values that already reference stored media are kept after they are - checked to be inside this server's media storage. + checked to be inside this server's media storage. Media URLs and ``url`` + pieces are downloaded once into the results store; a ``url`` piece takes + the path data type of its content. Args: pieces (Sequence[MessagePieceRequest]): Request pieces to resolve in place. @@ -312,26 +407,32 @@ async def persist_message_pieces_async( """ for piece in pieces: original_value = piece.original_value + original_type = piece.data_type converted_value = piece.converted_value converted_type = piece.converted_value_data_type or piece.data_type - if piece.data_type in MEDIA_PATH_DATA_TYPES: + mirrors_original = converted_value is None or ( + converted_value == piece.original_value and converted_type == piece.data_type + ) + original_resolved = False + if original_type in _IMPORTED_DATA_TYPES: result = await persist_media_value_async( value=original_value, - data_type=piece.data_type, + data_type=original_type, mime_type=piece.mime_type, serializer_factory=serializer_factory, ) if result.resolved: + original_resolved = True original_value = result.value - if converted_value is None or ( - converted_value == piece.original_value and converted_type == piece.data_type - ): + original_type = result.data_type or original_type + if mirrors_original: converted_value = original_value + converted_type = original_type if ( converted_value is not None - and converted_type in MEDIA_PATH_DATA_TYPES - and (converted_value != original_value or converted_type != piece.data_type) + and converted_type in _IMPORTED_DATA_TYPES + and not (mirrors_original and original_resolved) ): result = await persist_media_value_async( value=converted_value, @@ -340,6 +441,10 @@ async def persist_message_pieces_async( ) if result.resolved: converted_value = result.value + converted_type = result.data_type or converted_type piece.original_value = original_value + piece.data_type = original_type piece.converted_value = converted_value + if piece.converted_value_data_type is not None or converted_type != original_type: + piece.converted_value_data_type = converted_type diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py new file mode 100644 index 0000000000..ceeb170972 --- /dev/null +++ b/pyrit/backend/services/media_url_import.py @@ -0,0 +1,172 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Bounded download of caller-supplied media URLs into managed storage.""" + +from __future__ import annotations + +import asyncio +import mimetypes +import re +from dataclasses import dataclass +from pathlib import PurePosixPath +from urllib.parse import urlparse + +import httpx + +from pyrit.common.net_utility import get_httpx_client + +MAX_MEDIA_URL_BYTES = 100 * 1024 * 1024 +MAX_MEDIA_URL_REDIRECTS = 3 +_TIMEOUT = httpx.Timeout(30.0, connect=10.0) +_DEADLINE_SECONDS = 60.0 +_GENERIC_CONTENT_TYPES = frozenset({"application/octet-stream", "binary/octet-stream"}) +_URL_SUFFIX_PATTERN = re.compile(r"^\.[A-Za-z0-9]{1,10}$") + +_url_import_enabled = True + + +@dataclass(frozen=True) +class MediaDownload: + """Bytes downloaded from a media URL and what the server reported about them.""" + + content: bytes + content_type: str | None + final_url: str + + +def set_media_url_import_enabled(*, enabled: bool) -> None: + """Turn the download of caller-supplied media URLs on or off.""" + global _url_import_enabled + _url_import_enabled = enabled + + +def redact_url(url: str) -> str: + """ + Return a URL without credentials, query string, or fragment, for use in error messages. + + Returns: + str: The scheme, host, port, and path of the URL. + """ + parsed = urlparse(url) + host = parsed.hostname or "" + if parsed.port: + host = f"{host}:{parsed.port}" + return f"{parsed.scheme}://{host}{parsed.path}" + + +def media_content_type(download: MediaDownload) -> str | None: + """ + Return the reported media type, or the type implied by the URL suffix when it is missing or generic. + + Returns: + str | None: The lowercase media type without parameters, or None when unknown. + """ + content_type = (download.content_type or "").split(";", 1)[0].strip().lower() + if content_type and content_type not in _GENERIC_CONTENT_TYPES: + return content_type + guessed, _ = mimetypes.guess_type(urlparse(download.final_url).path, strict=False) + return guessed + + +def media_extension(download: MediaDownload, *, default: str) -> str: + """ + Choose the file extension for downloaded media. + + Returns: + str: The extension implied by the media type, else a short suffix from the URL path, else ``default``. + """ + content_type = media_content_type(download) + extension = mimetypes.guess_extension(content_type, strict=False) if content_type else None + if extension: + return extension + suffix = PurePosixPath(urlparse(download.final_url).path).suffix + return suffix.lower() if _URL_SUFFIX_PATTERN.match(suffix) else default + + +def _create_client() -> httpx.AsyncClient: + return get_httpx_client(use_async=True, timeout=_TIMEOUT, headers={"Accept-Encoding": "identity"}) + + +def _is_plain_http_url(url: httpx.URL) -> bool: + """ + Return whether a URL is http(s), names a host, and carries no user information. + + Returns: + bool: True for a URL the download may request. + """ + return url.scheme in ("http", "https") and bool(url.host) and not url.userinfo + + +async def _read_download_async(*, response: httpx.Response, shown: str) -> MediaDownload: + """ + Read a final response body within the size limit. + + Returns: + MediaDownload: The downloaded bytes and the reported content type. + + Raises: + ValueError: If the body is larger than the limit. + httpx.HTTPStatusError: If the response is not successful. + """ + response.raise_for_status() + too_large = f"Media URL {shown} is larger than {MAX_MEDIA_URL_BYTES // (1024 * 1024)} MiB." + declared_length = response.headers.get("content-length", "") + if declared_length.isdigit() and int(declared_length) > MAX_MEDIA_URL_BYTES: + raise ValueError(too_large) + chunks: list[bytes] = [] + total = 0 + async for chunk in response.aiter_bytes(): + total += len(chunk) + if total > MAX_MEDIA_URL_BYTES: + raise ValueError(too_large) + chunks.append(chunk) + return MediaDownload( + content=b"".join(chunks), + content_type=response.headers.get("content-type"), + final_url=str(response.url), + ) + + +async def download_media_url_async(*, url: str) -> MediaDownload: + """ + Download a media URL once, with a time limit, a size limit, and a redirect limit. + + No credentials or headers from the API request are forwarded. + + Returns: + MediaDownload: The downloaded bytes and the reported content type. + + Raises: + ValueError: If URL import is turned off, the URL is not a plain http(s) URL, or the download fails or + exceeds a limit. + """ + if not _url_import_enabled: + raise ValueError("Media URL import is turned off on this server; upload the file content instead.") + parsed = urlparse(url) + if parsed.scheme not in ("http", "https") or not parsed.hostname: + raise ValueError("Media URLs must be http or https URLs.") + if parsed.username is not None or parsed.password is not None: + raise ValueError("Media URLs must not include credentials.") + + shown = redact_url(url) + try: + async with asyncio.timeout(_DEADLINE_SECONDS), _create_client() as client: + request = client.build_request("GET", url) + # Redirects are followed here rather than by httpx, so each hop is checked and a + # redirect body is closed unread instead of being buffered outside the size limit. + for _ in range(MAX_MEDIA_URL_REDIRECTS + 1): + response = await client.send(request, stream=True, follow_redirects=False) + try: + if response.next_request is None: + return await _read_download_async(response=response, shown=shown) + request = response.next_request + finally: + await response.aclose() + if not _is_plain_http_url(request.url): + raise ValueError(f"Media URL {shown} redirected to a URL that is not a plain http or https URL.") + raise ValueError(f"Media URL {shown} redirected more than {MAX_MEDIA_URL_REDIRECTS} times.") + except httpx.HTTPStatusError as exc: + raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from exc + except (httpx.HTTPError, TimeoutError) as exc: + raise ValueError(f"Media URL {shown} could not be downloaded.") from exc diff --git a/pyrit/backend/services/runtime_lifecycle.py b/pyrit/backend/services/runtime_lifecycle.py index 202ff2cb3f..af2c21f062 100644 --- a/pyrit/backend/services/runtime_lifecycle.py +++ b/pyrit/backend/services/runtime_lifecycle.py @@ -14,6 +14,7 @@ from pyrit.backend.models.initializers import ConfiguredInitializerSetting from pyrit.backend.services.configuration_file_service import ConfigurationFileService from pyrit.backend.services.environment_file_service import EnvironmentFileService +from pyrit.backend.services.media_url_import import set_media_url_import_enabled from pyrit.backend.services.scenario_run_service import get_scenario_run_service, peek_scenario_run_service from pyrit.backend.services.service_lifecycle import close_services_async, outstanding_estimates from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH @@ -86,6 +87,7 @@ async def _management_async(self, config: ConfigurationLoader) -> None: read_only_file_sources=read_only, ) self.app.state.allow_custom_initializers = config.allow_custom_initializers + set_media_url_import_enabled(enabled=config.allow_media_url_import) registry = await asyncio.to_thread(InitializerRegistry.get_registry_singleton) registry.configure_custom_scripts_source(config.custom_initializers_source) diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index d821144c64..3b58b23af3 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -117,6 +117,8 @@ class ConfigurationLoader(YamlLoadable): operation: Name for the current operation. enable_live_reinitialization: Whether administrators may replace the live single-process backend runtime from the GUI. + allow_media_url_import: Whether the backend downloads http(s) media URLs from API + requests into managed storage. When False, such URLs are rejected. Example YAML configuration: memory_db_type: sqlite @@ -161,6 +163,7 @@ class ConfigurationLoader(YamlLoadable): max_concurrent_scenario_runs: int = 3 enable_live_reinitialization: bool = False allow_custom_initializers: bool = False + allow_media_url_import: bool = True custom_initializers_source: str | None = None server: dict[str, Any] | None = None extensions: dict[str, Any] = field(default_factory=dict) @@ -183,6 +186,8 @@ def __post_init__(self) -> None: validate_env_akv_strict(env_akv_strict=self.env_akv_strict) if not isinstance(self.enable_live_reinitialization, bool): raise TypeError("enable_live_reinitialization must be a bool.") + if not isinstance(self.allow_media_url_import, bool): + raise TypeError("allow_media_url_import must be a bool.") self._validate_allow_custom_initializers() self._normalize_memory_db_type() self._normalize_initializers() diff --git a/tests/unit/backend/test_attack_models.py b/tests/unit/backend/test_attack_models.py index b4c89dc883..49f9d1966e 100644 --- a/tests/unit/backend/test_attack_models.py +++ b/tests/unit/backend/test_attack_models.py @@ -8,7 +8,7 @@ import pytest from pydantic import ValidationError -from pyrit.backend.models.attacks import AddMessageRequest, MessagePieceRequest +from pyrit.backend.models.attacks import MessagePieceRequest from pyrit.models import PromptDataType @@ -74,25 +74,3 @@ def test_message_piece_accepts_applied_converter_order(converter_ids: list[str]) def test_message_piece_rejects_applied_converters_without_value(converter_ids: list[str]) -> None: with pytest.raises(ValidationError, match="applied_converter_ids requires converted_value"): MessagePieceRequest(original_value="source", applied_converter_ids=converter_ids) - - -@pytest.mark.parametrize( - "piece", - [ - MessagePieceRequest(data_type="url", original_value="https://example.test/image.png"), - MessagePieceRequest( - original_value="source", converted_value="https://example.test/image.png", converted_value_data_type="url" - ), - ], -) -def test_add_message_rejects_url_pieces_when_sending(piece: MessagePieceRequest) -> None: - with pytest.raises(ValidationError, match="URL pieces cannot be sent"): - AddMessageRequest(pieces=[piece], send=True, target_conversation_id="conversation") - - -def test_add_message_stores_url_pieces_without_sending() -> None: - piece = MessagePieceRequest(data_type="url", original_value="https://example.test/blob.png") - - request = AddMessageRequest(role="assistant", pieces=[piece], send=False, target_conversation_id="conversation") - - assert request.pieces[0].data_type == "url" diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index f50c18cf05..140b3926ab 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -25,6 +25,7 @@ ConverterService, get_converter_service, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.converter import ( Base64Converter, BinaryConverter, @@ -628,22 +629,62 @@ async def test_create_with_path_or_str_url( f"{_BLOB_RESULTS_ROOT}/private/input.mp4", f"{_BLOB_RESULTS_ROOT}/prompt-memory-entries/inputs/..\\..\\private/input.mp4", f"{_BLOB_RESULTS_ROOT}%5Cprompt-memory-entries/private/input.mp4", + "https://example.org/input.mp4", ], ) - async def test_path_or_str_url_outside_results_is_rejected( + async def test_path_or_str_url_outside_results_is_downloaded_to_owned_upload( self, upload_service: ConverterService, url: str ) -> None: - with _results_root(_BLOB_RESULTS_ROOT), pytest.raises(ValueError, match="result storage"): + download = MediaDownload(content=b"MP4", content_type="video/mp4", final_url=url) + with ( + _results_root(_BLOB_RESULTS_ROOT), + patch( + "pyrit.backend.services.converter_service.download_media_url_async", AsyncMock(return_value=download) + ) as download_mock, + ): + result, owned_paths = await upload_service._persist_data_uri_params_async( + converter_type="AddImageVideoConverter", params={"video_path": url} + ) + + download_mock.assert_awaited_once_with(url=url) + assert len(owned_paths) == 1 + assert owned_paths[0].parent == upload_service._upload_path + assert owned_paths[0].suffix == ".mp4" + assert owned_paths[0].read_bytes() == b"MP4" + assert result == {"video_path": owned_paths[0]} + + async def test_path_or_str_blob_url_is_kept_without_query(self, upload_service: ConverterService) -> None: + url = f"{_BLOB_RESULTS_ROOT}/prompt-memory-entries/inputs/input.mp4" + + with _results_root(_BLOB_RESULTS_ROOT): + result, owned_paths = await upload_service._persist_data_uri_params_async( + converter_type="AddImageVideoConverter", params={"video_path": f"{url}?sv=2024&sig=secret"} + ) + + assert result == {"video_path": url} + assert owned_paths == [] + + async def test_failed_url_download_leaves_no_upload(self, upload_service: ConverterService) -> None: + with ( + patch( + "pyrit.backend.services.converter_service.download_media_url_async", + AsyncMock(side_effect=ValueError("Media URL https://example.org/x.png returned HTTP 404.")), + ), + pytest.raises(ValueError, match="HTTP 404"), + ): await upload_service.create_converter_async( - request=CreateConverterRequest(name="remote", type="AddImageVideoConverter", params={"video_path": url}) + request=CreateConverterRequest( + name="remote", type="AddImageVideoConverter", params={"video_path": "https://example.org/x.png"} + ) ) assert upload_service._registry.instances.get_entry("remote") is None + assert list(upload_service._upload_path.iterdir()) == [] - @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( + @pytest.mark.parametrize("value", [r"C:\server\input.mp4", "input.mp4", "ftp://example.org/input.mp4", 123]) + async def test_path_or_str_rest_rejects_non_upload_non_url_values( self, upload_service: ConverterService, value: object ) -> None: - with pytest.raises(ValueError, match="data URI or supplied as an Azure Blob URL"): + with pytest.raises(ValueError, match="data URI or given as an http\\(s\\) URL"): await upload_service.create_converter_async( request=CreateConverterRequest( name="invalid", type="AddImageVideoConverter", params={"video_path": value} @@ -979,21 +1020,42 @@ async def test_preview_conversion_resolves_media_reference_without_persistence(s assert result.converted_value == str(stored_path) factory.assert_not_called() - @pytest.mark.parametrize("value", ["https://example.test/image.png", "/api/media?path=/etc/hostname"]) - async def test_preview_conversion_rejects_media_outside_results(self, tmp_path: Path, value: str) -> None: - """Media values must point into this server's media storage.""" + async def test_preview_conversion_rejects_media_outside_results(self, tmp_path: Path) -> None: + """Media references must point into this server's media storage.""" service = ConverterService() - request = ConverterPreviewRequest(original_value=value, original_value_data_type="image_path", converter_ids=[]) + request = ConverterPreviewRequest( + original_value="/api/media?path=/etc/hostname", original_value_data_type="image_path", converter_ids=[] + ) - with _results_root(str(tmp_path)), pytest.raises(ValueError, match="result storage|results directory"): + with _results_root(str(tmp_path)), pytest.raises(ValueError, match="results directory"): await service.preview_conversion_async(request=request) - def test_preview_request_rejects_url_input(self) -> None: - """Preview never downloads caller-supplied URLs.""" - with pytest.raises(ValidationError, match="URL input is not supported"): - ConverterPreviewRequest( - original_value="https://example.test/image.png", original_value_data_type="url", converter_ids=[] - ) + async def test_preview_imports_url_input_as_media(self, tmp_path: Path) -> None: + """A URL preview input is downloaded once and converted as the stored media type.""" + service = ConverterService() + serializer = MagicMock(value=str(tmp_path / "prompt-memory-entries" / "imported.png")) + serializer.save_data_async = AsyncMock() + download = MediaDownload(content=b"PNG", content_type="image/png", final_url="https://example.test/cat") + request = ConverterPreviewRequest( + original_value="https://example.test/cat", original_value_data_type="url", converter_ids=[] + ) + + with ( + _results_root(str(tmp_path)), + patch( + "pyrit.backend.services.media_persistence.download_media_url_async", AsyncMock(return_value=download) + ) as download_mock, + patch( + "pyrit.backend.services.converter_service.data_serializer_factory", return_value=serializer + ) as factory, + ): + result = await service.preview_conversion_async(request=request) + + download_mock.assert_awaited_once_with(url="https://example.test/cat") + factory.assert_called_once_with(category="prompt-memory-entries", data_type="image_path", extension=".png") + serializer.save_data_async.assert_awaited_once_with(b"PNG") + assert result.original_value_data_type == "url" + assert (result.converted_value, result.converted_value_data_type) == (serializer.value, "image_path") async def test_preview_conversion_chains_multiple_converters(self) -> None: """Test that preview chains multiple converters.""" diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index 202c913a8d..7a6e0d2545 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -9,12 +9,15 @@ import pytest +from pyrit.backend.models.attacks import MessagePieceRequest from pyrit.backend.services.media_persistence import ( MEDIA_SUBDIRECTORIES, MediaOrigin, persist_media_value_async, + persist_message_pieces_async, require_managed_blob_url, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.memory import CentralMemory from pyrit.memory.storage.storage import AzureBlobStorageIO @@ -25,6 +28,7 @@ def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: serializer = MagicMock() serializer.value = value serializer.save_b64_image_async = AsyncMock() + serializer.save_data_async = AsyncMock() return serializer @@ -32,6 +36,11 @@ def _results_root(root: str): return patch.object(CentralMemory, "get_memory_instance", return_value=MagicMock(results_path=root)) +def _download(*, content_type: str | None, final_url: str = "https://example.test/media", content: bytes = b"MEDIA"): + download = MediaDownload(content=content, content_type=content_type, final_url=final_url) + return patch("pyrit.backend.services.media_persistence.download_media_url_async", AsyncMock(return_value=download)) + + @pytest.fixture def stored_image(tmp_path: Path) -> Path: image_path = tmp_path / "results" / "prompt-memory-entries" / "images" / "stored.png" @@ -70,39 +79,38 @@ async def test_existing_local_path_inside_results_is_not_persisted(stored_image: factory.assert_not_called() -async def test_blob_url_inside_results_container_is_kept() -> None: +async def test_blob_url_inside_results_container_is_kept_without_query() -> None: value = f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png?sv=2024&sig=signature" with _results_root(_BLOB_ROOT): result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) assert result.origin is MediaOrigin.REMOTE_URL - assert result.value == value + assert result.value == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" assert result.persisted is False @pytest.mark.parametrize( "value", [ - "https://localHoSt/", - "https://127.0.1.2/", - "https://0177.0.23.19/", - "https://2130706433/", - "https://0x7f.00331.0246.174/", - "https://[::1]/", - "https://[fc00::]/", - "https://169.254.169.254/", "https://example.test/media.png", "https://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", ], ) -async def test_url_is_rejected_when_results_are_local(stored_image: Path, value: str) -> None: - factory = MagicMock() +async def test_url_is_downloaded_into_managed_storage_when_results_are_local(stored_image: Path, value: str) -> None: + serializer = _serializer(value="/results/prompt-memory-entries/images/imported.png") + factory = MagicMock(return_value=serializer) - with _results_root(str(stored_image.parents[2])), pytest.raises(ValueError, match="result storage"): - await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + with _results_root(str(stored_image.parents[2])), _download(content_type="image/png") as download: + result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) - factory.assert_not_called() + download.assert_awaited_once_with(url=value) + factory.assert_called_once_with(category="prompt-memory-entries", data_type="image_path", extension=".png") + serializer.save_data_async.assert_awaited_once_with(b"MEDIA") + assert result.origin is MediaOrigin.REMOTE_URL + assert result.value == "/results/prompt-memory-entries/images/imported.png" + assert result.persisted is True + assert result.data_type == "image_path" @pytest.mark.parametrize( @@ -125,9 +133,98 @@ async def test_url_is_rejected_when_results_are_local(stored_image: Path, value: f"{_BLOB_ROOT}/prompt-memory-entries", ], ) -async def test_url_outside_results_container_is_rejected(value: str) -> None: - with _results_root(_BLOB_ROOT), pytest.raises(ValueError, match="result storage"): - await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) +async def test_url_outside_managed_media_is_downloaded_not_read_from_storage(value: str) -> None: + """URLs outside the managed media folders are fetched like any URL, never read with the server's storage access.""" + factory = MagicMock(return_value=_serializer()) + + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + + download.assert_awaited_once_with(url=value) + assert result.persisted is True + assert result.value == "/saved/media.bin" + + +@pytest.mark.parametrize( + ("content_type", "final_url", "expected_type", "expected_extension"), + [ + ("image/png", "https://example.test/a", "image_path", ".png"), + ("audio/wav", "https://example.test/a", "audio_path", ".wav"), + ("video/mp4", "https://example.test/a", "video_path", ".mp4"), + ("application/pdf", "https://example.test/a", "binary_path", ".pdf"), + ("application/octet-stream", "https://example.test/a/photo.jpg", "image_path", ".jpg"), + (None, "https://example.test/a/blob", "binary_path", ".bin"), + ], +) +async def test_url_piece_takes_the_data_type_of_its_content( + content_type: str | None, final_url: str, expected_type: str, expected_extension: str +) -> None: + factory = MagicMock(return_value=_serializer()) + + with _results_root(_BLOB_ROOT), _download(content_type=content_type, final_url=final_url): + result = await persist_media_value_async( + value="https://example.test/a", data_type="url", serializer_factory=factory + ) + + factory.assert_called_once_with( + category="prompt-memory-entries", data_type=expected_type, extension=expected_extension + ) + assert result.data_type == expected_type + + +@pytest.mark.parametrize( + ("data_type", "content_type"), + [ + ("image_path", "text/html"), + ("image_path", "audio/wav"), + ("audio_path", "image/png"), + ("video_path", "text/plain"), + ], +) +async def test_downloaded_content_of_another_media_family_is_rejected(data_type: str, content_type: str) -> None: + factory = MagicMock() + + with ( + _results_root(_BLOB_ROOT), + _download(content_type=content_type), + pytest.raises(ValueError, match="not .* content") as error, + ): + await persist_media_value_async( + value="https://example.test/a.png?sig=secret", data_type=data_type, serializer_factory=factory + ) + + assert "secret" not in str(error.value) + factory.assert_not_called() + + +@pytest.mark.parametrize("content_type", ["application/octet-stream", "application/ogg", None]) +async def test_generic_content_is_stored_under_the_declared_type(content_type: str | None) -> None: + factory = MagicMock(return_value=_serializer()) + + with _results_root(_BLOB_ROOT), _download(content_type=content_type, final_url="https://example.test/a"): + result = await persist_media_value_async( + value="https://example.test/a", data_type="audio_path", serializer_factory=factory + ) + + assert result.data_type == "audio_path" + + +@pytest.mark.parametrize("value", ["example.test/image.png", "data:image/png;base64,UE5H", "/api/media?path=x"]) +async def test_url_piece_requires_an_http_url(value: str) -> None: + with _results_root(_BLOB_ROOT), pytest.raises(ValueError, match="http or https"): + await persist_media_value_async(value=value, data_type="url", serializer_factory=MagicMock()) + + +async def test_managed_blob_url_piece_is_kept_and_typed_by_extension() -> None: + value = f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png?sig=signature" + + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + result = await persist_media_value_async(value=value, data_type="url", serializer_factory=MagicMock()) + + download.assert_not_awaited() + assert result.value == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" + assert result.data_type == "image_path" + assert result.persisted is False @pytest.mark.parametrize( @@ -304,3 +401,44 @@ async def test_persistence_failure_returns_no_partial_result() -> None: data_type="binary_path", serializer_factory=MagicMock(return_value=serializer), ) + + +async def test_url_piece_is_imported_and_retyped_with_its_mirrored_conversion() -> None: + piece = MessagePieceRequest(data_type="url", original_value="https://example.test/cat") + factory = MagicMock(return_value=_serializer(value="/stored/cat.png")) + + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) + + download.assert_awaited_once() + assert (piece.data_type, piece.original_value) == ("image_path", "/stored/cat.png") + assert piece.converted_value == "/stored/cat.png" + assert piece.converted_value_data_type is None + + +async def test_converted_url_is_imported_separately_from_text_original() -> None: + piece = MessagePieceRequest( + original_value="describe this", converted_value="https://example.test/cat", converted_value_data_type="url" + ) + factory = MagicMock(return_value=_serializer(value="/stored/cat.png")) + + with _results_root(_BLOB_ROOT), _download(content_type="image/png"): + await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) + + assert (piece.data_type, piece.original_value) == ("text", "describe this") + assert (piece.converted_value_data_type, piece.converted_value) == ("image_path", "/stored/cat.png") + + +async def test_media_url_piece_keeps_its_declared_type() -> None: + piece = MessagePieceRequest(data_type="audio_path", original_value="https://example.test/clip.wav") + factory = MagicMock(return_value=_serializer(value="/stored/clip.wav")) + + with _results_root(_BLOB_ROOT), _download(content_type="audio/wav"): + await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) + + assert (piece.data_type, piece.original_value, piece.converted_value) == ( + "audio_path", + "/stored/clip.wav", + "/stored/clip.wav", + ) + assert piece.converted_value_data_type is None diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py new file mode 100644 index 0000000000..e9677d84be --- /dev/null +++ b/tests/unit/backend/test_media_url_import.py @@ -0,0 +1,206 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the bounded media URL download.""" + +from collections.abc import Callable, Iterator + +import httpx +import pytest + +from pyrit.backend.services import media_url_import +from pyrit.backend.services.media_url_import import ( + MAX_MEDIA_URL_REDIRECTS, + MediaDownload, + download_media_url_async, + media_content_type, + media_extension, + redact_url, + set_media_url_import_enabled, +) + +Handler = Callable[[httpx.Request], httpx.Response] + + +@pytest.fixture +def transport(monkeypatch: pytest.MonkeyPatch) -> Callable[[Handler], list[httpx.Request]]: + """Route media downloads through a mock transport and record the requests it receives.""" + + def install(handler: Handler) -> list[httpx.Request]: + requests: list[httpx.Request] = [] + + def record(request: httpx.Request) -> httpx.Response: + requests.append(request) + return handler(request) + + def create_client() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(record), headers={"Accept-Encoding": "identity"}) + + monkeypatch.setattr(media_url_import, "_create_client", create_client) + return requests + + return install + + +@pytest.fixture(autouse=True) +def url_import_enabled() -> Iterator[None]: + yield + set_media_url_import_enabled(enabled=True) + + +async def test_download_returns_content_and_type(transport: Callable[[Handler], list[httpx.Request]]) -> None: + requests = transport(lambda request: httpx.Response(200, content=b"PNG", headers={"content-type": "image/png"})) + + download = await download_media_url_async(url="https://example.test/cat.png?sig=secret") + + assert download == MediaDownload( + content=b"PNG", content_type="image/png", final_url="https://example.test/cat.png?sig=secret" + ) + assert len(requests) == 1 + assert requests[0].headers["accept-encoding"] == "identity" + assert "authorization" not in requests[0].headers + + +async def test_declared_length_over_limit_is_rejected(transport: Callable[[Handler], list[httpx.Request]]) -> None: + transport( + lambda request: httpx.Response( + 200, content=b"x", headers={"content-length": str(media_url_import.MAX_MEDIA_URL_BYTES + 1)} + ) + ) + + with pytest.raises(ValueError, match="larger than 100 MiB"): + await download_media_url_async(url="https://example.test/big.bin") + + +async def test_streamed_body_over_limit_is_rejected( + transport: Callable[[Handler], list[httpx.Request]], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(media_url_import, "MAX_MEDIA_URL_BYTES", 4) + + async def body(): + for chunk in (b"abc", b"def"): + yield chunk + + transport(lambda request: httpx.Response(200, content=body())) + + with pytest.raises(ValueError, match="larger than"): + await download_media_url_async(url="https://example.test/stream.bin") + + +async def test_too_many_redirects_are_rejected(transport: Callable[[Handler], list[httpx.Request]]) -> None: + requests = transport( + lambda request: httpx.Response(302, headers={"location": f"https://example.test/{len(request.url.path)}x"}) + ) + + with pytest.raises(ValueError, match=f"more than {MAX_MEDIA_URL_REDIRECTS} times"): + await download_media_url_async(url="https://example.test/start") + assert len(requests) == MAX_MEDIA_URL_REDIRECTS + 1 + + +@pytest.mark.parametrize("location", ["file:///etc/hostname", "https://user:secret@example.test/next.png"]) +async def test_redirect_to_non_plain_http_url_is_rejected( + transport: Callable[[Handler], list[httpx.Request]], location: str +) -> None: + requests = transport(lambda request: httpx.Response(302, headers={"location": location})) + + with pytest.raises(ValueError, match="redirected to a URL that is not a plain http or https URL") as error: + await download_media_url_async(url="https://example.test/start") + assert len(requests) == 1 + assert "secret" not in str(error.value) + + +async def test_redirect_body_is_not_read( + transport: Callable[[Handler], list[httpx.Request]], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(media_url_import, "MAX_MEDIA_URL_BYTES", 4) + read_redirect_body = False + + async def redirect_body(): + nonlocal read_redirect_body + read_redirect_body = True + yield b"x" * 1024 + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/start": + return httpx.Response(302, headers={"location": "/done.png"}, content=redirect_body()) + return httpx.Response(200, content=b"ok", headers={"content-type": "image/png"}) + + transport(handler) + + download = await download_media_url_async(url="https://example.test/start") + + assert download.content == b"ok" + assert download.final_url == "https://example.test/done.png" + assert not read_redirect_body + + +async def test_error_status_is_rejected_without_query_string( + transport: Callable[[Handler], list[httpx.Request]], +) -> None: + transport(lambda request: httpx.Response(404)) + + with pytest.raises(ValueError, match="returned HTTP 404") as error: + await download_media_url_async(url="https://user.example.test/cat.png?sv=1&sig=secret#frag") + assert "secret" not in str(error.value) + assert "https://user.example.test/cat.png" in str(error.value) + + +async def test_network_error_is_rejected(transport: Callable[[Handler], list[httpx.Request]]) -> None: + def fail(request: httpx.Request) -> httpx.Response: + raise httpx.ReadTimeout("slow", request=request) + + transport(fail) + + with pytest.raises(ValueError, match="could not be downloaded"): + await download_media_url_async(url="https://example.test/slow.png") + + +@pytest.mark.parametrize( + ("url", "message"), + [ + ("ftp://example.test/cat.png", "http or https"), + ("https:///cat.png", "http or https"), + ("https://user:pass@example.test/cat.png", "credentials"), + ], +) +async def test_invalid_urls_are_rejected_before_download( + transport: Callable[[Handler], list[httpx.Request]], url: str, message: str +) -> None: + requests = transport(lambda request: httpx.Response(200)) + + with pytest.raises(ValueError, match=message): + await download_media_url_async(url=url) + assert requests == [] + + +async def test_disabled_import_rejects_urls(transport: Callable[[Handler], list[httpx.Request]]) -> None: + requests = transport(lambda request: httpx.Response(200)) + set_media_url_import_enabled(enabled=False) + + with pytest.raises(ValueError, match="turned off"): + await download_media_url_async(url="https://example.test/cat.png") + assert requests == [] + + +def test_redact_url_drops_credentials_query_and_fragment() -> None: + assert redact_url("https://user:pw@example.test:8443/a/b.png?sig=1#x") == "https://example.test:8443/a/b.png" + + +@pytest.mark.parametrize( + ("content_type", "final_url", "expected_type", "expected_extension"), + [ + ("image/png; charset=binary", "https://example.test/a", "image/png", ".png"), + ("application/octet-stream", "https://example.test/a/photo.JPG", "image/jpeg", ".jpg"), + (None, "https://example.test/a/clip.mp4?sig=1", "video/mp4", ".mp4"), + (None, "https://example.test/a/blob", None, ".bin"), + ("application/octet-stream", "https://example.test/a/file.weird123", None, ".weird123"), + (None, "https://example.test/a/file.toolongsuffix1", None, ".bin"), + ], +) +def test_media_type_and_extension( + content_type: str | None, final_url: str, expected_type: str | None, expected_extension: str +) -> None: + download = MediaDownload(content=b"", content_type=content_type, final_url=final_url) + + assert media_content_type(download) == expected_type + assert media_extension(download, default=".bin") == expected_extension diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index e4894b2d94..1c466d2ea5 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -27,6 +27,7 @@ ManualSendScheduler, get_manual_send_scheduler, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.backend.services.message_send_service import MessageSendService, resolve_applied_converter_identifiers from pyrit.backend.services.target_service import TargetService from pyrit.common.attack_result_scope import get_current_attack_result_id @@ -795,7 +796,8 @@ async def test_converted_media_references_are_not_repersisted_async( stored_path = (tmp_path / "prompt-memory-entries" / "preview.png").resolve() if stored_in_blob: mock_memory.results_path = "https://account.blob.core.windows.net/results" - converted_value = expected_value = f"{mock_memory.results_path}/prompt-memory-entries/preview.png?sv=1" + expected_value = f"{mock_memory.results_path}/prompt-memory-entries/preview.png" + converted_value = f"{expected_value}?sv=1" else: mock_memory.results_path = str(tmp_path) converted_value, expected_value = f"/api/media?path={stored_path}", str(stored_path) @@ -817,16 +819,15 @@ async def test_converted_media_references_are_not_repersisted_async( assert request.pieces[0].converted_value == expected_value factory.assert_not_called() - @pytest.mark.parametrize("converted_value", ["https://example.com/preview.png", "/api/media?path=/etc/hostname"]) async def test_converted_media_outside_results_is_rejected_async( - self, *, mock_memory: MagicMock, tmp_path: Path, converted_value: str + self, *, mock_memory: MagicMock, tmp_path: Path ) -> None: mock_memory.results_path = str(tmp_path) request = AddMessageRequest( pieces=[ MessagePieceRequest( original_value="source", - converted_value=converted_value, + converted_value="/api/media?path=/etc/hostname", converted_value_data_type="image_path", ) ], @@ -834,9 +835,33 @@ async def test_converted_media_outside_results_is_rejected_async( target_conversation_id="test-id", ) - with pytest.raises(ValueError, match="result storage|results directory"): + with pytest.raises(ValueError, match="results directory"): await MessageSendService._persist_base64_pieces_async(request) + async def test_url_piece_is_imported_before_sending_async(self, *, tmp_path: Path) -> None: + stored = str(tmp_path / "prompt-memory-entries" / "imported.png") + serializer = MagicMock(value=stored) + serializer.save_data_async = AsyncMock() + download = MediaDownload(content=b"PNG", content_type="image/png", final_url="https://example.com/cat") + request = AddMessageRequest( + pieces=[MessagePieceRequest(data_type="url", original_value="https://example.com/cat?sig=secret")], + send=True, + target_conversation_id="test-id", + ) + + with ( + patch( + "pyrit.backend.services.media_persistence.download_media_url_async", AsyncMock(return_value=download) + ) as download_mock, + patch("pyrit.backend.services.message_send_service.data_serializer_factory", return_value=serializer), + ): + await MessageSendService._persist_base64_pieces_async(request) + + download_mock.assert_awaited_once_with(url="https://example.com/cat?sig=secret") + piece = request.pieces[0] + assert (piece.data_type, piece.original_value) == ("image_path", stored) + assert "sig=secret" not in piece.model_dump_json() + async def test_identical_original_and_converted_media_saved_once_async(self) -> None: request = AddMessageRequest( pieces=[ @@ -1084,16 +1109,16 @@ async def test_path_data_type_supplies_extension_when_mime_type_missing(self, me ) assert request.pieces[0].original_value == "/saved/image.png" - async def test_http_url_is_kept_as_is(self, message_send_service, mock_memory: MagicMock) -> None: - """HTTPS blob URLs inside the results container should not be re-persisted.""" + async def test_http_url_is_kept_without_query(self, message_send_service, mock_memory: MagicMock) -> None: + """HTTPS blob URLs inside the results container are kept as references, without their query string.""" mock_memory.results_path = "https://myblob.blob.core.windows.net/results" - blob_url = "https://myblob.blob.core.windows.net/results/prompt-memory-entries/images/photo.png?sv=2024" + blob_url = "https://myblob.blob.core.windows.net/results/prompt-memory-entries/images/photo.png" request = AddMessageRequest( role="user", pieces=[ MessagePieceRequest( data_type="image_path", - original_value=blob_url, + original_value=f"{blob_url}?sv=2024&sig=secret", mime_type="image/png", ), ], diff --git a/tests/unit/backend/test_runtime_lifecycle.py b/tests/unit/backend/test_runtime_lifecycle.py index 87d5489b33..f6fbd6bdd8 100644 --- a/tests/unit/backend/test_runtime_lifecycle.py +++ b/tests/unit/backend/test_runtime_lifecycle.py @@ -257,6 +257,19 @@ async def test_invalid_cold_configuration_requires_restart_but_retains_raw_repai assert "in_memory" in await source.read_async() +@pytest.mark.parametrize("enabled", [True, False]) +async def test_management_applies_media_url_import_setting(enabled: bool) -> None: + service = RuntimeLifecycle(app=FastAPI(), source=ConfigurationFileService(config_file_value=None)) + config = ConfigurationLoader(memory_db_type="in_memory", env_files=[], allow_media_url_import=enabled) + with ( + patch.object(lifecycle_module, "set_media_url_import_enabled") as set_enabled, + patch.object(lifecycle_module.InitializerRegistry, "get_registry_singleton", return_value=MagicMock()), + ): + await service._management_async(config) + + set_enabled.assert_called_once_with(enabled=enabled) + + def test_authorization_policy_does_not_change_with_environment(runtime: RuntimeLifecycle) -> None: runtime.app.state.auth_environment["PYRIT_ALLOW_UNAUTHENTICATED_ADMIN"] = "" request = Request({"type": "http", "app": runtime.app}) diff --git a/tests/unit/setup/test_configuration_loader.py b/tests/unit/setup/test_configuration_loader.py index 04b0b3f39d..7bee2b160c 100644 --- a/tests/unit/setup/test_configuration_loader.py +++ b/tests/unit/setup/test_configuration_loader.py @@ -76,6 +76,15 @@ def test_rejects_non_boolean_allow_custom_initializers(self, invalid_value: obje with pytest.raises(TypeError, match=r"allow_custom_initializers must be a bool"): ConfigurationLoader(allow_custom_initializers=invalid_value) # type: ignore[arg-type] + def test_media_url_import_is_enabled_by_default_and_can_be_disabled(self) -> None: + assert ConfigurationLoader().allow_media_url_import is True + assert ConfigurationLoader.from_dict({"allow_media_url_import": False}).allow_media_url_import is False + + @pytest.mark.parametrize("invalid_value", ["false", 0, 1, [], {}]) + def test_rejects_non_boolean_allow_media_url_import(self, invalid_value: object) -> None: + with pytest.raises(TypeError, match=r"allow_media_url_import must be a bool"): + ConfigurationLoader(allow_media_url_import=invalid_value) # type: ignore[arg-type] + def test_valid_memory_db_types_snake_case(self): """Test all valid memory database types in snake_case.""" for db_type in ["in_memory", "sqlite", "azure_sql"]: From 2e58dc6cbf0b2a36b44bd317b34a8c0b114bb2e7 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 17:48:23 +0000 Subject: [PATCH 3/7] Keep media URLs as references unless imported, and separate code loading from file access - Keep http(s) media values and url pieces as references by default. A caller imports a URL by setting import_url on a piece or preview, and the stored copy keeps the declared media path type. - Reject Azure Blob URLs outside the result storage as references, since the storage layer would read them from the server's own container. - Return the stored copy and its source metadata from a preview that imports, and record the redacted source URL and content type on imported pieces. - Name the reason and limit in download errors, and log each failure with the classes in its cause chain. - Replace uses_host_resources with loads_local_code and upload_directory_parameter: code-loading targets stay out of the API, HTTPXAPITarget uploads only from the operator's target_upload_directory, and target parameters that name server paths cannot be set through the API. --- .pyrit_conf_example | 16 +- pyrit/backend/README.md | 33 ++- pyrit/backend/models/attacks.py | 26 ++ pyrit/backend/models/converters.py | 36 ++- pyrit/backend/services/attack_service.py | 5 +- pyrit/backend/services/converter_service.py | 21 +- pyrit/backend/services/media_persistence.py | 155 +++++++---- pyrit/backend/services/media_url_import.py | 46 +++- pyrit/backend/services/runtime_lifecycle.py | 2 + pyrit/backend/services/target_service.py | 97 +++++-- pyrit/prompt_target/common/prompt_target.py | 12 +- .../http_target/httpx_api_target.py | 3 +- .../hugging_face/hugging_face_chat_target.py | 2 +- pyrit/setup/configuration_loader.py | 10 +- tests/unit/backend/test_attack_service.py | 43 +++ tests/unit/backend/test_converter_service.py | 47 +++- tests/unit/backend/test_media_persistence.py | 247 +++++++++++------- tests/unit/backend/test_media_url_import.py | 67 ++++- .../unit/backend/test_message_send_service.py | 8 +- tests/unit/backend/test_runtime_lifecycle.py | 13 + tests/unit/backend/test_target_service.py | 48 +++- tests/unit/setup/test_configuration_loader.py | 10 + 22 files changed, 716 insertions(+), 231 deletions(-) diff --git a/.pyrit_conf_example b/.pyrit_conf_example index b797d8fb25..2ac9b927d5 100644 --- a/.pyrit_conf_example +++ b/.pyrit_conf_example @@ -134,14 +134,22 @@ enable_live_reinitialization: false # Default: false allow_custom_initializers: false -# When true, the backend downloads http(s) media URLs from API requests (media -# values, `url` pieces, and converter file parameters) once into managed storage, -# with time, size, and redirect limits, and passes only the stored copy on. When -# false, such URLs are rejected and media must be uploaded. +# When true, the backend downloads http(s) media URLs that API callers ask it to +# import (`import_url` on message pieces and converter previews, and URLs given for +# converter file parameters) once into managed storage, with time, size, and redirect +# limits, and passes only the stored copy on. Media URLs that are not imported stay +# references. When false, import requests are rejected and media must be uploaded. # # Default: true allow_media_url_import: true +# Directory that targets created through the backend API may upload local files from +# (HTTPXAPITarget). The server passes it to the target; API callers cannot choose it. +# When unset, such targets cannot be created through the API. +# +# Default: unset +# target_upload_directory: /path/to/upload/files + # Optional storage for custom initializer Python scripts. This may be a local # directory or an Azure Blob container URI with an optional blob prefix. # Container URIs may include a SAS; otherwise DefaultAzureCredential is used. Defaults to diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index e4a1dbce42..eca59a81d8 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -210,19 +210,26 @@ request values before using them: converters, must be uploaded content, a media URL, or a reference into this server's media storage: the `prompt-memory-entries` and `seed-prompt-entries` folders under the memory results path. Other file paths are rejected. -- Media URLs and `url` pieces are downloaded once into that storage (60 second limit, - 100 MiB limit, at most 3 redirects, no request credentials forwarded), and only the stored - copy reaches converters and targets, so signed URLs are not passed to model providers. - This includes messages that are stored without being sent, because stored history is - replayed to targets later. A `url` piece takes the media type of its content - (`image_path`, `audio_path`, `video_path`, or `binary_path`); send a literal URL as `text`. - Blob URLs inside the configured results container are kept as references, without their - query string, instead of being downloaded. Set `allow_media_url_import: false` in - `.pyrit_conf` to reject media URLs instead. -- Target types that read local files or load model code (`HTTPXAPITarget`, - `HuggingFaceChatTarget`) cannot be created through the API. Register them in Python or - with an initializer, where the operator controls their settings; for example, set - `HTTPXAPITarget(allowed_upload_directory=...)` so uploads stay inside one folder. +- Media URLs are kept as references by default, and `url` pieces always pass through + unchanged. Azure Blob URLs outside the configured results container are rejected as media + references, because the storage layer would read them from this server's own container; + blob URLs inside it are kept without their query string. +- A caller imports a media URL by setting `import_url` on a message piece or a converter + preview with an `image_path`, `audio_path`, `video_path`, or `binary_path` type. The server + downloads it once into managed storage (10 second connect, 30 second read, and 60 second + total limits, 100 MiB limit, at most 3 redirects, no request credentials forwarded) and + stores it under the declared type. Converters and targets then see only the stored copy, + and the piece's prompt metadata records the source URL, without credentials or query + string, and the reported content type. A preview that imports returns the stored copy and + that metadata, so sending them reuses the same bytes. Set `allow_media_url_import: false` + in `.pyrit_conf` to turn imports off. Converter file parameters given a URL are downloaded + the same way. +- Target types that load model code (`HuggingFaceChatTarget`) cannot be created through the + API, and target parameters that name server paths cannot be set through it. Targets that + upload local files (`HTTPXAPITarget`) can be created through the API only when + `target_upload_directory` is set in `.pyrit_conf`; the server passes that directory to the + target, which uploads files only from inside it. Register such targets in Python or with an + initializer for other settings. Intentional exceptions: diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index 3b280c09d6..763161b6d1 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -23,6 +23,7 @@ TextStr, ) from pyrit.models import ( + MEDIA_PATH_DATA_TYPES, AttackResult, ChatMessageRole, ConversationReference, @@ -398,6 +399,31 @@ class MessagePieceRequest(BaseModel): source_piece_id: uuid.UUID | None = Field( None, description="Source piece for a complete conversation save; verified against its source conversation." ) + import_url: bool = Field( + False, + description="Download this piece's http(s) media values once into managed storage instead of keeping the " + "URLs as references. The stored copies keep the declared image_path, audio_path, video_path, or " + "binary_path type.", + ) + + @model_validator(mode="after") + def _validate_import_url(self) -> "MessagePieceRequest": + """ + Validate that a URL import names the media type to store. + + Returns: + The validated request piece. + + Raises: + ValueError: If import_url is set on a piece without a media path type. + """ + converted_type = self.converted_value_data_type or self.data_type + if self.import_url and not {self.data_type, converted_type} & MEDIA_PATH_DATA_TYPES: + raise ValueError( + "import_url needs an image_path, audio_path, video_path, or binary_path value; " + "declare the media type to store the URL as." + ) + return self @model_validator(mode="after") def _validate_converted_value_data_type(self) -> "MessagePieceRequest": diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 28c5a46f40..9bfff4f5df 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -9,10 +9,10 @@ from typing import Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from pyrit.backend.models.common import MAX_ITEMS, REGISTRY_INSTANCE_NAME_PATTERN, IdentifierStr -from pyrit.models import ConverterIdentifier, Parameter, PromptDataType +from pyrit.models import MEDIA_PATH_DATA_TYPES, ConverterIdentifier, Parameter, PromptDataType __all__ = [ "ConverterInstance", @@ -121,13 +121,43 @@ class ConverterPreviewRequest(BaseModel): converter_ids: list[IdentifierStr] = Field(..., max_length=MAX_ITEMS, description="Converter instance IDs to apply") start_token: str = Field(default="⟪", min_length=1, description="Opening marker for selected text regions") end_token: str = Field(default="⟫", min_length=1, description="Closing marker for selected text regions") + import_url: bool = Field( + False, + description="Download an http(s) original_value once into managed storage instead of passing the URL on. " + "The stored copy keeps the declared image_path, audio_path, video_path, or binary_path type.", + ) + + @model_validator(mode="after") + def _validate_import_url(self) -> "ConverterPreviewRequest": + """ + Validate that a URL import names the media type to store. + + Returns: + The validated preview request. + + Raises: + ValueError: If import_url is set without a media path type. + """ + if self.import_url and self.original_value_data_type not in MEDIA_PATH_DATA_TYPES: + raise ValueError( + "import_url needs an image_path, audio_path, video_path, or binary_path original_value_data_type; " + "declare the media type to store the URL as." + ) + return self class ConverterPreviewResponse(BaseModel): """Response from converter preview.""" - original_value: str = Field(..., description="Original input text") + original_value: str = Field( + ..., description="Original input, or the stored copy that conversion used when an http(s) URL was imported" + ) original_value_data_type: PromptDataType = Field(..., description="Data type of original value") converted_value: str = Field(..., description="Final converted text") converted_value_data_type: PromptDataType = Field(..., description="Data type of converted value") steps: list[PreviewStep] = Field(..., description="Step-by-step conversion results") + prompt_metadata: dict[str, str] = Field( + default_factory=dict, + description="Source information for an imported original value. Send it as the piece's prompt_metadata " + "with the stored copy to keep it in the conversation.", + ) diff --git a/pyrit/backend/services/attack_service.py b/pyrit/backend/services/attack_service.py index fcaf9ae1e5..00632e58fa 100644 --- a/pyrit/backend/services/attack_service.py +++ b/pyrit/backend/services/attack_service.py @@ -54,7 +54,7 @@ ) from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.message_sends import MessageSendRequest, MessageSendStatus -from pyrit.backend.services.media_persistence import persist_message_pieces_async +from pyrit.backend.services.media_persistence import media_source_entries, persist_message_pieces_async from pyrit.backend.services.message_send_service import ( MessageSendService, get_message_send_service, @@ -785,10 +785,9 @@ async def _prepare_message_pieces_async( for saved, request_piece in prepared_pieces: await self._persist_base64_pieces_async(pieces=[request_piece], persisted_paths=persisted_paths) saved.original_value = request_piece.original_value - saved.original_value_data_type = request_piece.data_type converted_value = request_piece.converted_value saved.converted_value = converted_value if converted_value is not None else saved.original_value - saved.converted_value_data_type = request_piece.converted_value_data_type or request_piece.data_type + saved.prompt_metadata.update(media_source_entries(request_piece.prompt_metadata)) await set_message_piece_sha256_async(saved) return [saved for saved, _ in prepared_pieces] diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index 4f20bb2e7c..0be70023b6 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -37,7 +37,11 @@ CreateConverterRequest, PreviewStep, ) -from pyrit.backend.services.media_persistence import is_managed_blob_url, persist_media_value_async +from pyrit.backend.services.media_persistence import ( + is_managed_blob_url, + media_source_metadata, + persist_media_value_async, +) from pyrit.backend.services.media_url_import import download_media_url_async, media_extension from pyrit.common.azure_storage import redact_url_credentials from pyrit.memory import data_serializer_factory @@ -217,16 +221,19 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> For non-text data types (image_path, audio_path, etc.), persists base64 data to a temporary file so converters can operate on file paths. Marked text - regions use the request's delimiter settings for every stage. + regions use the request's delimiter settings for every stage. When the + request imports an http(s) URL, the response returns the stored copy and + its source metadata, so a later send can reuse the same bytes. Returns: ConverterPreviewResponse with step-by-step conversion results. """ original_value = request.original_value data_type = request.original_value_data_type + source_metadata: dict[str, str] = {} - # For path-based data types and URLs, resolve references, import URLs, or persist base64/data URIs. - if str(data_type).endswith("_path") or data_type == "url": + # For path-based data types, resolve references, import URLs on request, or persist base64/data URIs. + if str(data_type).endswith("_path"): result = await persist_media_value_async( value=original_value, data_type=data_type, @@ -235,10 +242,11 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> # explicit/data-URI MIME metadata. use_data_uri_mime_type=False, require_valid_base64_after_path_error=True, + import_url=request.import_url, serializer_factory=data_serializer_factory, ) original_value = result.value - data_type = result.data_type or data_type + source_metadata = media_source_metadata(result) converters = self._gather_converters(converter_ids=request.converter_ids) steps, final_value, final_type = await self._apply_converters_async( @@ -250,11 +258,12 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> ) return ConverterPreviewResponse( - original_value=request.original_value, + original_value=original_value if source_metadata else request.original_value, original_value_data_type=request.original_value_data_type, converted_value=final_value, converted_value_data_type=final_type, steps=steps, + prompt_metadata=source_metadata, ) def get_converter_objects_for_ids(self, *, converter_ids: list[str]) -> list[Any]: diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index fb94a28a5c..29c2f35fd6 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -9,7 +9,7 @@ import base64 import binascii import mimetypes -from collections.abc import Callable, Coroutine, Sequence +from collections.abc import Callable, Coroutine, Mapping, Sequence from dataclasses import dataclass from enum import Enum from pathlib import Path @@ -18,9 +18,9 @@ from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS from pyrit.backend.services.media_url_import import ( + MediaDownload, download_media_url_async, media_content_type, - media_extension, redact_url, ) from pyrit.common.azure_storage import is_azure_blob_uri, redact_url_credentials @@ -36,8 +36,11 @@ # Media type families of the path data types; binary_path accepts any content. _MEDIA_FAMILIES: dict[PromptDataType, str] = {"image_path": "image", "audio_path": "audio", "video_path": "video"} _CHECKED_FAMILIES = frozenset({"image", "audio", "video", "text"}) -# Data types whose request values are stored as managed media before use. -_IMPORTED_DATA_TYPES = frozenset({*MEDIA_PATH_DATA_TYPES, "url"}) +# Prompt metadata that records where imported media came from, for the original and converted values. +_MEDIA_SOURCE_PREFIXES = ("media_source", "converted_media_source") +MEDIA_SOURCE_METADATA_KEYS = frozenset( + f"{prefix}_{field}" for prefix in _MEDIA_SOURCE_PREFIXES for field in ("url", "content_type") +) class MediaOrigin(str, Enum): @@ -60,7 +63,7 @@ class MediaPersistenceResult: resolved: bool mime_type: str | None = None extension: str | None = None - data_type: PromptDataType | None = None + source: str | None = None SerializerFactory = Callable[..., Any] @@ -251,18 +254,21 @@ async def _write_owned_media_async( raise -def _downloaded_data_type(*, content_type: str | None) -> PromptDataType: +def _imported_extension(*, download: MediaDownload, data_type: PromptDataType) -> str: """ - Return the path data type for downloaded ``url`` content. + Choose the file extension for imported media of a declared path type. Returns: - PromptDataType: ``image_path``, ``audio_path``, or ``video_path`` by media family, else ``binary_path``. + str: The extension of the reported (or URL-implied) media type when it fits the declared + type, else the declared type's default extension. """ - family = (content_type or "").split("/", 1)[0] - for data_type, data_family in _MEDIA_FAMILIES.items(): - if family == data_family: - return data_type - return "binary_path" + content_type = media_content_type(download) + expected_family = _MEDIA_FAMILIES.get(data_type) + if content_type and expected_family in (None, content_type.split("/", 1)[0]): + extension = mimetypes.guess_extension(content_type, strict=False) + if extension: + return extension + return DEFAULT_MEDIA_EXTENSIONS.get(str(data_type), ".bin") async def _import_media_url_async( @@ -273,28 +279,27 @@ async def _import_media_url_async( created_paths: list[str] | None, ) -> MediaPersistenceResult: """ - Download a media URL once and store the bytes in managed media storage. + Download a media URL once and store the bytes in managed media storage under the declared type. - A ``url`` value is stored under the path data type of its content. Other path types keep - their declared type and must not receive content of a different media family. + A missing or generic reported content type is accepted, since the caller declared the type; + content of a different media family than the declared type is rejected. Returns: - MediaPersistenceResult: The managed reference to the stored copy. + MediaPersistenceResult: The managed reference to the stored copy and its redacted source URL. Raises: ValueError: If the download fails or the content does not match the declared media type. """ download = await download_media_url_async(url=url) content_type = media_content_type(download) - resolved_type = _downloaded_data_type(content_type=content_type) if data_type == "url" else data_type - expected_family = _MEDIA_FAMILIES.get(resolved_type) + expected_family = _MEDIA_FAMILIES.get(data_type) family = (content_type or "").split("/", 1)[0] if expected_family and family in _CHECKED_FAMILIES and family != expected_family: raise ValueError(f"Media URL {redact_url(url)} returned {content_type}, not {expected_family} content.") - extension = media_extension(download, default=DEFAULT_MEDIA_EXTENSIONS.get(str(resolved_type), ".bin")) + extension = _imported_extension(download=download, data_type=data_type) serializer = serializer_factory( category="prompt-memory-entries", - data_type=resolved_type, + data_type=data_type, extension=extension, ) await _write_owned_media_async( @@ -309,10 +314,23 @@ async def _import_media_url_async( resolved=True, mime_type=content_type, extension=extension, - data_type=resolved_type, + source=redact_url(url), ) +def _is_read_from_result_storage(url: str) -> bool: + """ + Return whether the storage layer reads this URL from the server's own blob container. + + The storage layer treats any URL on a ``blob.core.windows.net`` host as Azure storage and reads + it from the configured container, keeping only the blob path. + + Returns: + bool: True when reading the URL would read a blob in this server's result storage. + """ + return "blob.core.windows.net" in urlparse(url).netloc + + async def persist_media_value_async( *, value: str, @@ -320,11 +338,12 @@ async def persist_media_value_async( mime_type: str | None = None, use_data_uri_mime_type: bool = True, require_valid_base64_after_path_error: bool = False, + import_url: bool = False, serializer_factory: SerializerFactory = data_serializer_factory, created_paths: list[str] | None = None, ) -> MediaPersistenceResult: """ - Classify and, when needed, persist one path-typed or ``url`` media value. + Classify and, when needed, persist one path-typed media value. The two policy flags preserve the small historical differences between attack ingestion and converter preview while keeping origin detection, @@ -332,33 +351,44 @@ async def persist_media_value_async( ``/api/media`` references are accepted only when they point into this server's media storage, so callers cannot make the server read other files. Blob URLs inside this server's result storage are kept as references, without - their query string, since the server reads them with its own credentials; any - other http(s) URL is downloaded once into managed storage, so converters and - targets only see the stored copy. + their query string, since the server reads them with its own credentials. + Other http(s) URLs are kept as references, or downloaded once into managed + storage under the declared type when ``import_url`` is set; Azure Blob URLs + outside the result storage must be imported, because the storage layer would + read them from this server's own container. Returns: A typed result containing the resolved value and persistence metadata. Raises: ValueError: If the value names a file outside the server's media storage, or a - URL cannot be imported. + URL cannot be imported or must be imported. """ if value.startswith(("http://", "https://")): if is_managed_blob_url(value): - blob_type, _ = mimetypes.guess_type(urlparse(value).path, strict=False) return MediaPersistenceResult( value=redact_url_credentials(value), origin=MediaOrigin.REMOTE_URL, persisted=False, resolved=True, mime_type=mime_type, - data_type=_downloaded_data_type(content_type=blob_type) if data_type == "url" else data_type, ) - return await _import_media_url_async( - url=value, data_type=data_type, serializer_factory=serializer_factory, created_paths=created_paths + if import_url: + return await _import_media_url_async( + url=value, data_type=data_type, serializer_factory=serializer_factory, created_paths=created_paths + ) + if _is_read_from_result_storage(value): + raise ValueError( + f"Media URL {redact_url(value)} points to Azure Blob Storage outside this server's result storage. " + "Set import_url to download it, or send it as a url piece." + ) + return MediaPersistenceResult( + value=value, + origin=MediaOrigin.REMOTE_URL, + persisted=False, + resolved=True, + mime_type=mime_type, ) - if data_type == "url": - raise ValueError("URL pieces must use an http or https URL.") reference_path = _media_reference_path(value) if reference_path is not None: @@ -420,6 +450,36 @@ async def persist_media_value_async( ) +def media_source_entries(metadata: Mapping[str, Any] | None) -> dict[str, Any]: + """ + Return the media source entries that persistence recorded in a piece's prompt metadata. + + Returns: + dict[str, Any]: The entries whose keys are in ``MEDIA_SOURCE_METADATA_KEYS``. + """ + return {key: value for key, value in (metadata or {}).items() if key in MEDIA_SOURCE_METADATA_KEYS} + + +def media_source_metadata(result: MediaPersistenceResult, *, prefix: str = "media_source") -> dict[str, str]: + """ + Return prompt metadata naming where an imported media value came from. + + Args: + result (MediaPersistenceResult): The persistence result for one value. + prefix (str): Key prefix, so a piece can record its original and converted sources apart. + + Returns: + dict[str, str]: The redacted source URL and the content type the server reported, or an + empty dict when the value was not imported. + """ + if result.source is None: + return {} + metadata = {f"{prefix}_url": result.source} + if result.mime_type: + metadata[f"{prefix}_content_type"] = result.mime_type + return metadata + + async def persist_message_pieces_async( *, pieces: Sequence[MessagePieceRequest], @@ -434,9 +494,10 @@ async def persist_message_pieces_async( values to be **file paths**, so base64 data is written to the results store and the request values are replaced with the resulting file path. Values that already reference stored media are kept after they are - checked to be inside this server's media storage. Media URLs and ``url`` - pieces are downloaded once into the results store; a ``url`` piece takes - the path data type of its content. + checked to be inside this server's media storage. Media URLs stay + references unless the piece sets ``import_url``; imported values are + stored under their declared type, and the piece's prompt metadata records + where they came from. ``url`` pieces are left unchanged. Args: pieces (Sequence[MessagePieceRequest]): Request pieces to resolve in place. @@ -446,46 +507,48 @@ async def persist_message_pieces_async( """ for piece in pieces: original_value = piece.original_value - original_type = piece.data_type converted_value = piece.converted_value converted_type = piece.converted_value_data_type or piece.data_type mirrors_original = converted_value is None or ( converted_value == piece.original_value and converted_type == piece.data_type ) + source_metadata: dict[str, str] = {} original_resolved = False - if original_type in _IMPORTED_DATA_TYPES: + if piece.data_type in MEDIA_PATH_DATA_TYPES: result = await persist_media_value_async( value=original_value, - data_type=original_type, + data_type=piece.data_type, mime_type=piece.mime_type, + import_url=piece.import_url, serializer_factory=serializer_factory, created_paths=persisted_paths, ) if result.resolved: original_resolved = True original_value = result.value - original_type = result.data_type or original_type + source_metadata.update(media_source_metadata(result)) if mirrors_original: converted_value = original_value - converted_type = original_type if ( converted_value is not None - and converted_type in _IMPORTED_DATA_TYPES + and converted_type in MEDIA_PATH_DATA_TYPES and not (mirrors_original and original_resolved) ): result = await persist_media_value_async( value=converted_value, data_type=converted_type, + import_url=piece.import_url, serializer_factory=serializer_factory, created_paths=persisted_paths, ) if result.resolved: converted_value = result.value - converted_type = result.data_type or converted_type + source_metadata.update(media_source_metadata(result, prefix="converted_media_source")) piece.original_value = original_value - piece.data_type = original_type piece.converted_value = converted_value - if piece.converted_value_data_type is not None or converted_type != original_type: - piece.converted_value_data_type = converted_type + if source_metadata: + # Without other metadata the message mapper records the piece's mime_type; keep it. + existing = piece.prompt_metadata or ({"mime_type": piece.mime_type} if piece.mime_type else {}) + piece.prompt_metadata = {**existing, **source_metadata} diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py index ceeb170972..8d6a7b7e1b 100644 --- a/pyrit/backend/services/media_url_import.py +++ b/pyrit/backend/services/media_url_import.py @@ -6,6 +6,7 @@ from __future__ import annotations import asyncio +import logging import mimetypes import re from dataclasses import dataclass @@ -16,9 +17,13 @@ from pyrit.common.net_utility import get_httpx_client +logger = logging.getLogger(__name__) + MAX_MEDIA_URL_BYTES = 100 * 1024 * 1024 MAX_MEDIA_URL_REDIRECTS = 3 -_TIMEOUT = httpx.Timeout(30.0, connect=10.0) +_CONNECT_TIMEOUT_SECONDS = 10.0 +_READ_TIMEOUT_SECONDS = 30.0 +_TIMEOUT = httpx.Timeout(_READ_TIMEOUT_SECONDS, connect=_CONNECT_TIMEOUT_SECONDS) _DEADLINE_SECONDS = 60.0 _GENERIC_CONTENT_TYPES = frozenset({"application/octet-stream", "binary/octet-stream"}) _URL_SUFFIX_PATTERN = re.compile(r"^\.[A-Za-z0-9]{1,10}$") @@ -167,6 +172,43 @@ async def download_media_url_async(*, url: str) -> MediaDownload: raise ValueError(f"Media URL {shown} redirected to a URL that is not a plain http or https URL.") raise ValueError(f"Media URL {shown} redirected more than {MAX_MEDIA_URL_REDIRECTS} times.") except httpx.HTTPStatusError as exc: + _log_download_failure(shown=shown, reason=f"HTTP {exc.response.status_code}", exc=exc) raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from exc except (httpx.HTTPError, TimeoutError) as exc: - raise ValueError(f"Media URL {shown} could not be downloaded.") from exc + reason = _failure_reason(exc) + _log_download_failure(shown=shown, reason=reason, exc=exc) + raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from exc + + +def _failure_reason(exc: BaseException) -> str: + """ + Describe why a download failed, with the limit that was hit, without quoting the exception. + + httpx exception messages can include the full URL and its query string, so they are not repeated. + + Returns: + str: A short reason such as ``connecting timed out after 10 seconds``. + """ + if isinstance(exc, httpx.ConnectTimeout): + return f"connecting timed out after {_CONNECT_TIMEOUT_SECONDS:g} seconds" + if isinstance(exc, httpx.ReadTimeout): + return f"no data arrived for {_READ_TIMEOUT_SECONDS:g} seconds" + if isinstance(exc, httpx.WriteTimeout): + return f"sending the request timed out after {_READ_TIMEOUT_SECONDS:g} seconds" + if isinstance(exc, httpx.PoolTimeout): + return f"no connection was free within {_READ_TIMEOUT_SECONDS:g} seconds" + if isinstance(exc, TimeoutError): + return f"the download took longer than {_DEADLINE_SECONDS:g} seconds" + if isinstance(exc, httpx.ConnectError): + return "the connection failed" + return f"the request failed ({type(exc).__name__})" + + +def _log_download_failure(*, shown: str, reason: str, exc: BaseException) -> None: + """Log a failed download with the redacted URL and the exception classes in its cause chain.""" + causes: list[str] = [] + cause: BaseException | None = exc + while cause is not None and len(causes) < 4: + causes.append(type(cause).__name__) + cause = cause.__cause__ or cause.__context__ + logger.warning("Media URL %s could not be downloaded: %s (%s)", shown, reason, " <- ".join(causes)) diff --git a/pyrit/backend/services/runtime_lifecycle.py b/pyrit/backend/services/runtime_lifecycle.py index 95d742d0b6..94b8e44b58 100644 --- a/pyrit/backend/services/runtime_lifecycle.py +++ b/pyrit/backend/services/runtime_lifecycle.py @@ -21,6 +21,7 @@ has_active_manual_sends, outstanding_estimates, ) +from pyrit.backend.services.target_service import set_target_upload_directory from pyrit.common.path import CONFIGURATION_DIRECTORY_PATH from pyrit.memory import CentralMemory from pyrit.registry import InitializerRegistry @@ -103,6 +104,7 @@ async def _management_async(self, config: ConfigurationLoader) -> None: ) self.app.state.allow_custom_initializers = config.allow_custom_initializers set_media_url_import_enabled(enabled=config.allow_media_url_import) + set_target_upload_directory(directory=config.target_upload_directory) registry = await asyncio.to_thread(InitializerRegistry.get_registry_singleton) registry.configure_custom_scripts_source(config.custom_initializers_source) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 33520915f0..1d7a457d9e 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -15,6 +15,7 @@ import asyncio import logging import uuid +from collections.abc import Mapping, Sequence from functools import lru_cache from typing import Any, Literal @@ -29,10 +30,13 @@ from pyrit.common import REQUIRED_VALUE from pyrit.models.catalog.target import TargetInstance from pyrit.models.parameter import Parameter +from pyrit.prompt_target import PromptTarget from pyrit.registry import TargetRegistry logger = logging.getLogger(__name__) +_target_upload_directory: str | None = None + _ENV_BACKED_REQUIRED_PARAMETERS: dict[str, frozenset[str]] = { "OpenAITarget": frozenset({"endpoint", "model_name"}), "AzureBlobStorageTarget": frozenset({"container_url"}), @@ -43,6 +47,12 @@ } +def set_target_upload_directory(*, directory: str | None) -> None: + """Set the directory that targets created through the API may upload local files from.""" + global _target_upload_directory + _target_upload_directory = directory + + class TargetService: """ Service for managing target instances. @@ -184,16 +194,51 @@ def _project_target_parameters(self, *, target_type: str, parameters: tuple[Para for parameter in parameters ] + @staticmethod + def _reject_server_resources( + *, + target_type: str, + target_cls: type[PromptTarget], + params: Mapping[str, Any], + parameters: Sequence[Parameter], + ) -> None: + """ + Reject requests that would make a target load model code or read files chosen by the caller. + + Raises: + ValueError: If the target loads model code, uploads local files while no upload + directory is configured, or a supplied parameter names a server path. + """ + if target_cls.loads_local_code: + raise ValueError( + f"Target type '{target_type}' loads model code on this server and cannot be created through the " + "API. Register it in Python or with an initializer instead." + ) + if target_cls.upload_directory_parameter and _target_upload_directory is None: + raise ValueError( + f"Target type '{target_type}' uploads files from this server and can only be created through the " + "API when target_upload_directory is set in the server configuration. Register it in Python or " + "with an initializer instead." + ) + for parameter in parameters: + if parameter.name in params and (parameter.is_path or parameter.is_path_or_str): + raise ValueError( + f"Parameter '{parameter.name}' of '{target_type}' names a path on this server and cannot be " + "set through the API." + ) + async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. Returns every registered target with its derived constructor parameters and the auth modes it supports, all projected from the - registry's ``TargetMetadata``. Types that read local files or load model - code are listed but cannot be created through the API. Deciding which - entries to surface to a user is a presentation concern owned by the - caller (e.g. the frontend), not this service. + registry's ``TargetMetadata``. Types that load model code, or upload + local files when the server has no upload directory configured, are + listed but cannot be created through the API, and neither can parameters + that name server paths. Deciding which entries to surface to a user is a + presentation concern owned by the caller (e.g. the frontend), not this + service. Returns: TargetTypeResponse containing all available target classes. @@ -221,12 +266,13 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn reference resolution, and construction are owned by the ``TargetRegistry``. Endpoint trust and identity token minting are owned by the target classes themselves. This service only enforces the - request-level contract: it rejects target types that read local files or - load model code, and for ``identity`` it confirms the target supports it - and omits the api_key plus any registry-flagged identity-conflicting - parameters so the target validates its own endpoint and authenticates - itself. The response is built before the target is registered, so a - failed request leaves no registered target. + request-level contract: it rejects target types that load model code, + rejects parameters that name server paths, gives targets that upload local + files the operator-configured upload directory, and for ``identity`` it + confirms the target supports it and omits the api_key plus any + registry-flagged identity-conflicting parameters so the target validates + its own endpoint and authenticates itself. The response is built before + the target is registered, so a failed request leaves no registered target. Args: request: The create target request with type, params, and auth_mode. @@ -235,11 +281,12 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn TargetInstance with the new target's details. Raises: - ValueError: If the target type is not registered, uses local files or - model code, or identity auth is requested but unsupported by the - target type. Construction errors (unknown params, incompatible inner - targets, unrecognized identity endpoints) are raised by the - registry / target classes. + ValueError: If the target type is not registered, loads model code, + uploads local files without a configured upload directory, or a + parameter names a server path, or identity auth is requested but + unsupported by the target type. Construction errors (unknown params, + incompatible inner targets, unrecognized identity endpoints) are + raised by the registry / target classes. """ if request.type not in self._registry: raise ValueError( @@ -247,11 +294,11 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn ) target_cls = self._registry.get_class(request.type) - if target_cls.uses_host_resources: - raise ValueError( - f"Target type '{request.type}' reads local files or loads model code and cannot be created " - "through the API. Register it in Python or with an initializer instead." - ) + metadata = await asyncio.to_thread(self._registry.get_registered_class_metadata, request.type) + parameters = metadata.parameters if metadata is not None else () + self._reject_server_resources( + target_type=request.type, target_cls=target_cls, params=request.params, parameters=parameters + ) params: dict[str, Any] = dict(request.params) if request.auth_mode == "identity": @@ -262,12 +309,12 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn # Omit any other parameter the registry metadata marks as conflicting with # identity-based auth (e.g. AzureBlobStorageTarget's sas_token), so a caller # can't silently override the selected auth mode by also supplying it. - metadata = await asyncio.to_thread(self._registry.get_registered_class_metadata, request.type) - if metadata is not None: - for parameter in metadata.parameters: - if parameter.identity_conflicting: - params.pop(parameter.name, None) + for parameter in parameters: + if parameter.identity_conflicting: + params.pop(parameter.name, None) params.update(target_cls.get_auth_mode_parameters(auth_mode=request.auth_mode)) + if target_cls.upload_directory_parameter: + params[target_cls.upload_directory_parameter] = _target_upload_directory # LEGACY COMPATIBILITY: The current configuration UI omits the name. # Remove this generated fallback after that UI sends an explicit name. diff --git a/pyrit/prompt_target/common/prompt_target.py b/pyrit/prompt_target/common/prompt_target.py index fd06d3805c..11ae7eaabe 100644 --- a/pyrit/prompt_target/common/prompt_target.py +++ b/pyrit/prompt_target/common/prompt_target.py @@ -84,10 +84,14 @@ class PromptTarget(Identifiable): # Azure Blob Storage, Prompt Shield) override this to add ``"identity"``. supported_auth_modes: ClassVar[tuple[AuthMode, ...]] = ("api_key",) - # Declarative fact consumed by the create-target service. Targets that read files - # from the machine running PyRIT or load model code on it set this to True, and the - # create-target API rejects them; register those targets in Python or with an initializer. - uses_host_resources: ClassVar[bool] = False + # Declarative facts consumed by the create-target service. Targets that load model code + # on the machine running PyRIT set ``loads_local_code``, and the create-target API rejects + # them; register those targets in Python or with an initializer. Targets that upload files + # from that machine name the constructor parameter that confines those uploads in + # ``upload_directory_parameter``; the create-target API sets it to the directory the + # server operator configured, and rejects the type when none is configured. + loads_local_code: ClassVar[bool] = False + upload_directory_parameter: ClassVar[str | None] = None def __init_subclass__(cls, **kwargs: object) -> None: """ diff --git a/pyrit/prompt_target/http_target/httpx_api_target.py b/pyrit/prompt_target/http_target/httpx_api_target.py index 12bb912fbd..a5c90ae994 100644 --- a/pyrit/prompt_target/http_target/httpx_api_target.py +++ b/pyrit/prompt_target/http_target/httpx_api_target.py @@ -37,8 +37,7 @@ class HTTPXAPITarget(HTTPTarget): """ _PATH_TYPES: frozenset[str] = frozenset({"image_path", "audio_path", "video_path", "binary_path"}) - # Uploads files from the local file system. - uses_host_resources: ClassVar[bool] = True + upload_directory_parameter: ClassVar[str | None] = "allowed_upload_directory" _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( capabilities=TargetCapabilities( supports_multi_turn=True, diff --git a/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py b/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py index e42767f4b0..904c522425 100644 --- a/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py +++ b/pyrit/prompt_target/hugging_face/hugging_face_chat_target.py @@ -37,7 +37,7 @@ class HuggingFaceChatTarget(PromptTarget): ) # Loads models, and optionally their code, on the local machine. - uses_host_resources: ClassVar[bool] = True + loads_local_code: ClassVar[bool] = True # Class-level cache for model and tokenizer _cached_model: Any = None diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index 7d0047ff3b..bf6ab9dd90 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -118,8 +118,11 @@ class ConfigurationLoader(YamlLoadable): operation: Name for the current operation. enable_live_reinitialization: Whether administrators may replace the live single-process backend runtime from the GUI. - allow_media_url_import: Whether the backend downloads http(s) media URLs from API - requests into managed storage. When False, such URLs are rejected. + allow_media_url_import: Whether the backend downloads http(s) media URLs that API + callers ask it to import. When False, import requests are rejected. + target_upload_directory: Directory that targets created through the backend API may + upload local files from (``HTTPXAPITarget``). When unset, such targets cannot be + created through the API. Example YAML configuration: memory_db_type: sqlite @@ -165,6 +168,7 @@ class ConfigurationLoader(YamlLoadable): enable_live_reinitialization: bool = False allow_custom_initializers: bool = False allow_media_url_import: bool = True + target_upload_directory: str | None = None custom_initializers_source: str | None = None server: dict[str, Any] | None = None extensions: dict[str, Any] = field(default_factory=dict) @@ -189,6 +193,8 @@ def __post_init__(self) -> None: raise TypeError("enable_live_reinitialization must be a bool.") if not isinstance(self.allow_media_url_import, bool): raise TypeError("allow_media_url_import must be a bool.") + if self.target_upload_directory is not None and not isinstance(self.target_upload_directory, str): + raise TypeError("target_upload_directory must be a string path.") self._validate_allow_custom_initializers() self._normalize_memory_db_type() self._normalize_initializers() diff --git a/tests/unit/backend/test_attack_service.py b/tests/unit/backend/test_attack_service.py index 617ebeca85..313d81345f 100644 --- a/tests/unit/backend/test_attack_service.py +++ b/tests/unit/backend/test_attack_service.py @@ -44,6 +44,7 @@ ManualSendQueueFullError, ManualSendScheduler, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.backend.services.message_send_service import MessageSendService, get_message_send_service from pyrit.backend.services.pagination import ( decode_keyset_cursor, @@ -984,6 +985,48 @@ async def test_create_attack_persists_prepended_base64_media(self, attack_servic stored_piece = mock_memory.add_conversation_branches_to_attack_async.call_args.kwargs["message_pieces"][0] assert stored_piece.original_value == stored_path + async def test_create_attack_records_the_source_of_imported_prepended_media( + self, attack_service, mock_memory + ) -> None: + stored_path = "/results/prompt-memory-entries/images/imported.png" + serializer = MagicMock(value=stored_path) + serializer.get_data_filename_async = AsyncMock(return_value=stored_path) + serializer.save_data_async = AsyncMock() + download = MediaDownload(content=b"PNG", content_type="image/png", final_url="https://example.test/cat") + prepended = [ + PrependedMessageRequest( + role="user", + pieces=[ + MessagePieceRequest( + data_type="image_path", original_value="https://example.test/cat?sig=secret", import_url=True + ) + ], + ) + ] + + with ( + patch("pyrit.backend.services.attack_service.get_target_service") as mock_get_target_service, + patch("pyrit.backend.services.attack_service.data_serializer_factory", return_value=serializer), + patch( + "pyrit.backend.services.media_persistence.download_media_url_async", AsyncMock(return_value=download) + ), + patch("pyrit.backend.services.attack_service.set_message_piece_sha256_async", new_callable=AsyncMock), + ): + mock_target_service = MagicMock() + mock_target_service.get_target_async = AsyncMock(return_value=MagicMock(type="TextTarget")) + mock_target_service.get_target_object.return_value.get_identifier.return_value = ComponentIdentifier( + class_name="TextTarget", class_module="pyrit.prompt_target" + ) + mock_get_target_service.return_value = mock_target_service + await attack_service.create_attack_async( + request=CreateAttackRequest(target_registry_name="target-1", prepended_conversation=prepended) + ) + + stored_piece = mock_memory.add_conversation_branches_to_attack_async.call_args.kwargs["message_pieces"][0] + assert (stored_piece.original_value_data_type, stored_piece.original_value) == ("image_path", stored_path) + assert stored_piece.prompt_metadata["media_source_url"] == "https://example.test/cat" + assert "secret" not in str(stored_piece.prompt_metadata) + async def test_create_attack_lowers_system_prompt_to_system_message(self, attack_service, mock_memory) -> None: """Test that system_prompt is lowered to a single system-role message at sequence 0.""" with patch("pyrit.backend.services.attack_service.get_target_service") as mock_get_target_service: diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index a792581b53..be0c4164d4 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -1131,14 +1131,18 @@ async def test_preview_conversion_rejects_media_outside_results(self, tmp_path: with _results_root(str(tmp_path)), pytest.raises(ValueError, match="results directory"): await service.preview_conversion_async(request=request) - async def test_preview_imports_url_input_as_media(self, tmp_path: Path) -> None: - """A URL preview input is downloaded once and converted as the stored media type.""" + async def test_preview_imports_url_input_and_returns_the_stored_copy(self, tmp_path: Path) -> None: + """An imported URL is downloaded once, and the response returns the stored copy for a later send.""" service = ConverterService() - serializer = MagicMock(value=str(tmp_path / "prompt-memory-entries" / "imported.png")) + stored = str(tmp_path / "prompt-memory-entries" / "imported.png") + serializer = MagicMock(value=stored) serializer.save_data_async = AsyncMock() download = MediaDownload(content=b"PNG", content_type="image/png", final_url="https://example.test/cat") request = ConverterPreviewRequest( - original_value="https://example.test/cat", original_value_data_type="url", converter_ids=[] + original_value="https://example.test/cat?sig=secret", + original_value_data_type="image_path", + converter_ids=[], + import_url=True, ) with ( @@ -1152,11 +1156,40 @@ async def test_preview_imports_url_input_as_media(self, tmp_path: Path) -> None: ): result = await service.preview_conversion_async(request=request) - download_mock.assert_awaited_once_with(url="https://example.test/cat") + download_mock.assert_awaited_once_with(url="https://example.test/cat?sig=secret") factory.assert_called_once_with(category="prompt-memory-entries", data_type="image_path", extension=".png") serializer.save_data_async.assert_awaited_once_with(b"PNG") - assert result.original_value_data_type == "url" - assert (result.converted_value, result.converted_value_data_type) == (serializer.value, "image_path") + assert (result.original_value, result.original_value_data_type) == (stored, "image_path") + assert (result.converted_value, result.converted_value_data_type) == (stored, "image_path") + assert result.prompt_metadata == { + "media_source_url": "https://example.test/cat", + "media_source_content_type": "image/png", + } + + @pytest.mark.parametrize("data_type", ["url", "image_path"]) + async def test_preview_keeps_url_input_as_a_reference_without_import(self, tmp_path: Path, data_type: str) -> None: + service = ConverterService() + url = "https://example.test/cat.png?sig=signature" + request = ConverterPreviewRequest(original_value=url, original_value_data_type=data_type, converter_ids=[]) + + with ( + _results_root(str(tmp_path)), + patch("pyrit.backend.services.media_persistence.download_media_url_async") as download_mock, + ): + result = await service.preview_conversion_async(request=request) + + download_mock.assert_not_called() + assert (result.original_value, result.converted_value) == (url, url) + assert result.prompt_metadata == {} + + def test_preview_import_requires_a_declared_media_type(self) -> None: + with pytest.raises(ValidationError, match="declare the media type"): + ConverterPreviewRequest( + original_value="https://example.test/cat", + original_value_data_type="url", + converter_ids=[], + import_url=True, + ) async def test_preview_conversion_chains_multiple_converters(self) -> None: """Test that preview chains multiple converters.""" diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index 7a6e0d2545..125162badc 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -8,6 +8,7 @@ from urllib.parse import quote import pytest +from pydantic import ValidationError from pyrit.backend.models.attacks import MessagePieceRequest from pyrit.backend.services.media_persistence import ( @@ -93,16 +94,67 @@ async def test_blob_url_inside_results_container_is_kept_without_query() -> None @pytest.mark.parametrize( "value", [ - "https://example.test/media.png", - "https://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", + "https://example.test/media.png?sig=signature", + "http://example.test/media", + "https://169.254.169.254/results/prompt-memory-entries/images/stored.png", ], ) -async def test_url_is_downloaded_into_managed_storage_when_results_are_local(stored_image: Path, value: str) -> None: +async def test_url_is_kept_as_a_reference_without_import(value: str) -> None: + factory = MagicMock() + + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + + download.assert_not_awaited() + factory.assert_not_called() + assert result.origin is MediaOrigin.REMOTE_URL + assert result.value == value + assert result.persisted is False + assert result.source is None + + +_BLOB_URLS_OUTSIDE_MANAGED_MEDIA = [ + "https://other.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}-other/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}/other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/../other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/images/..\\..\\other-folder/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries/images/%5C..%5C..%5Cother-folder/stored.png", + "http://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}%2Fprompt-memory-entries/other-folder/stored.png", + f"{_BLOB_ROOT}%2fprompt-memory-entries/stored.png", + f"{_BLOB_ROOT}%5Cprompt-memory-entries/stored.png", + f"{_BLOB_ROOT}/./prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}//prompt-memory-entries/images/stored.png", + "https://account.blob.core.windows.net//results/prompt-memory-entries/images/stored.png", + f"{_BLOB_ROOT}/prompt-memory-entries", + "https://account.blob.core.windows.net.example.test/results/prompt-memory-entries/images/stored.png", +] + + +@pytest.mark.parametrize("value", _BLOB_URLS_OUTSIDE_MANAGED_MEDIA) +async def test_blob_url_outside_managed_media_is_rejected_without_import(value: str) -> None: + """The storage layer would read these from the server's own container, so they cannot stay references.""" + with ( + _results_root(_BLOB_ROOT), + _download(content_type="image/png") as download, + pytest.raises(ValueError, match="outside this server's result storage"), + ): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=MagicMock()) + + download.assert_not_awaited() + + +@pytest.mark.parametrize("value", ["https://example.test/media.png?sig=signature", *_BLOB_URLS_OUTSIDE_MANAGED_MEDIA]) +async def test_import_url_downloads_into_managed_storage(value: str) -> None: + """Imports are fetched like any URL, never read with the server's storage access.""" serializer = _serializer(value="/results/prompt-memory-entries/images/imported.png") factory = MagicMock(return_value=serializer) - with _results_root(str(stored_image.parents[2])), _download(content_type="image/png") as download: - result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + result = await persist_media_value_async( + value=value, data_type="image_path", import_url=True, serializer_factory=factory + ) download.assert_awaited_once_with(url=value) factory.assert_called_once_with(category="prompt-memory-entries", data_type="image_path", extension=".png") @@ -110,66 +162,45 @@ async def test_url_is_downloaded_into_managed_storage_when_results_are_local(sto assert result.origin is MediaOrigin.REMOTE_URL assert result.value == "/results/prompt-memory-entries/images/imported.png" assert result.persisted is True - assert result.data_type == "image_path" + assert result.source is not None + assert "?" not in result.source -@pytest.mark.parametrize( - "value", - [ - "https://other.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", - f"{_BLOB_ROOT}-other/prompt-memory-entries/images/stored.png", - f"{_BLOB_ROOT}/other-folder/stored.png", - f"{_BLOB_ROOT}/prompt-memory-entries/../other-folder/stored.png", - f"{_BLOB_ROOT}/prompt-memory-entries/images/..\\..\\other-folder/stored.png", - f"{_BLOB_ROOT}/prompt-memory-entries/images/%5C..%5C..%5Cother-folder/stored.png", - "http://account.blob.core.windows.net/results/prompt-memory-entries/images/stored.png", - "https://169.254.169.254/results/prompt-memory-entries/images/stored.png", - f"{_BLOB_ROOT}%2Fprompt-memory-entries/other-folder/stored.png", - f"{_BLOB_ROOT}%2fprompt-memory-entries/stored.png", - f"{_BLOB_ROOT}%5Cprompt-memory-entries/stored.png", - f"{_BLOB_ROOT}/./prompt-memory-entries/images/stored.png", - f"{_BLOB_ROOT}//prompt-memory-entries/images/stored.png", - "https://account.blob.core.windows.net//results/prompt-memory-entries/images/stored.png", - f"{_BLOB_ROOT}/prompt-memory-entries", - ], -) -async def test_url_outside_managed_media_is_downloaded_not_read_from_storage(value: str) -> None: - """URLs outside the managed media folders are fetched like any URL, never read with the server's storage access.""" - factory = MagicMock(return_value=_serializer()) +async def test_import_of_managed_blob_url_keeps_the_reference() -> None: + value = f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png?sig=signature" with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: - result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + result = await persist_media_value_async( + value=value, data_type="image_path", import_url=True, serializer_factory=MagicMock() + ) - download.assert_awaited_once_with(url=value) - assert result.persisted is True - assert result.value == "/saved/media.bin" + download.assert_not_awaited() + assert result.value == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" + assert result.persisted is False @pytest.mark.parametrize( - ("content_type", "final_url", "expected_type", "expected_extension"), + ("data_type", "content_type", "final_url", "expected_extension"), [ - ("image/png", "https://example.test/a", "image_path", ".png"), - ("audio/wav", "https://example.test/a", "audio_path", ".wav"), - ("video/mp4", "https://example.test/a", "video_path", ".mp4"), - ("application/pdf", "https://example.test/a", "binary_path", ".pdf"), - ("application/octet-stream", "https://example.test/a/photo.jpg", "image_path", ".jpg"), - (None, "https://example.test/a/blob", "binary_path", ".bin"), + ("image_path", "application/octet-stream", "https://example.test/a", ".png"), + ("image_path", None, "https://example.test/a", ".png"), + ("image_path", "application/octet-stream", "https://example.test/a/photo.jpg", ".jpg"), + ("audio_path", "application/ogg", "https://example.test/a", ".wav"), + ("binary_path", "application/pdf", "https://example.test/a", ".pdf"), + ("binary_path", "image/png", "https://example.test/a", ".png"), ], ) -async def test_url_piece_takes_the_data_type_of_its_content( - content_type: str | None, final_url: str, expected_type: str, expected_extension: str +async def test_import_stores_the_declared_type( + data_type: str, content_type: str | None, final_url: str, expected_extension: str ) -> None: factory = MagicMock(return_value=_serializer()) with _results_root(_BLOB_ROOT), _download(content_type=content_type, final_url=final_url): - result = await persist_media_value_async( - value="https://example.test/a", data_type="url", serializer_factory=factory + await persist_media_value_async( + value="https://example.test/a", data_type=data_type, import_url=True, serializer_factory=factory ) - factory.assert_called_once_with( - category="prompt-memory-entries", data_type=expected_type, extension=expected_extension - ) - assert result.data_type == expected_type + factory.assert_called_once_with(category="prompt-memory-entries", data_type=data_type, extension=expected_extension) @pytest.mark.parametrize( @@ -181,7 +212,7 @@ async def test_url_piece_takes_the_data_type_of_its_content( ("video_path", "text/plain"), ], ) -async def test_downloaded_content_of_another_media_family_is_rejected(data_type: str, content_type: str) -> None: +async def test_imported_content_of_another_media_family_is_rejected(data_type: str, content_type: str) -> None: factory = MagicMock() with ( @@ -190,43 +221,16 @@ async def test_downloaded_content_of_another_media_family_is_rejected(data_type: pytest.raises(ValueError, match="not .* content") as error, ): await persist_media_value_async( - value="https://example.test/a.png?sig=secret", data_type=data_type, serializer_factory=factory + value="https://example.test/a.png?sig=secret", + data_type=data_type, + import_url=True, + serializer_factory=factory, ) assert "secret" not in str(error.value) factory.assert_not_called() -@pytest.mark.parametrize("content_type", ["application/octet-stream", "application/ogg", None]) -async def test_generic_content_is_stored_under_the_declared_type(content_type: str | None) -> None: - factory = MagicMock(return_value=_serializer()) - - with _results_root(_BLOB_ROOT), _download(content_type=content_type, final_url="https://example.test/a"): - result = await persist_media_value_async( - value="https://example.test/a", data_type="audio_path", serializer_factory=factory - ) - - assert result.data_type == "audio_path" - - -@pytest.mark.parametrize("value", ["example.test/image.png", "data:image/png;base64,UE5H", "/api/media?path=x"]) -async def test_url_piece_requires_an_http_url(value: str) -> None: - with _results_root(_BLOB_ROOT), pytest.raises(ValueError, match="http or https"): - await persist_media_value_async(value=value, data_type="url", serializer_factory=MagicMock()) - - -async def test_managed_blob_url_piece_is_kept_and_typed_by_extension() -> None: - value = f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png?sig=signature" - - with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: - result = await persist_media_value_async(value=value, data_type="url", serializer_factory=MagicMock()) - - download.assert_not_awaited() - assert result.value == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" - assert result.data_type == "image_path" - assert result.persisted is False - - @pytest.mark.parametrize( "value", [ @@ -403,22 +407,74 @@ async def test_persistence_failure_returns_no_partial_result() -> None: ) -async def test_url_piece_is_imported_and_retyped_with_its_mirrored_conversion() -> None: - piece = MessagePieceRequest(data_type="url", original_value="https://example.test/cat") +async def test_url_piece_passes_through_unchanged() -> None: + piece = MessagePieceRequest(data_type="url", original_value="https://example.test/cat?sig=signature") + factory = MagicMock() + + with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: + await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) + + download.assert_not_awaited() + factory.assert_not_called() + assert (piece.data_type, piece.original_value) == ("url", "https://example.test/cat?sig=signature") + assert piece.converted_value is None + assert piece.prompt_metadata is None + + +async def test_imported_piece_keeps_its_type_and_records_its_source() -> None: + piece = MessagePieceRequest( + data_type="image_path", original_value="https://example.test/cat?sig=signature", import_url=True + ) factory = MagicMock(return_value=_serializer(value="/stored/cat.png")) with _results_root(_BLOB_ROOT), _download(content_type="image/png") as download: await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) download.assert_awaited_once() - assert (piece.data_type, piece.original_value) == ("image_path", "/stored/cat.png") - assert piece.converted_value == "/stored/cat.png" + assert (piece.data_type, piece.original_value, piece.converted_value) == ( + "image_path", + "/stored/cat.png", + "/stored/cat.png", + ) assert piece.converted_value_data_type is None + assert piece.prompt_metadata == { + "media_source_url": "https://example.test/cat", + "media_source_content_type": "image/png", + } + + +@pytest.mark.parametrize( + ("prompt_metadata", "mime_type", "kept"), + [ + ({"video_id": "abc"}, None, {"video_id": "abc"}), + (None, "image/png", {"mime_type": "image/png"}), + ], +) +async def test_imported_piece_keeps_existing_metadata( + prompt_metadata: dict[str, str] | None, mime_type: str | None, kept: dict[str, str] +) -> None: + piece = MessagePieceRequest( + data_type="image_path", + original_value="https://example.test/cat", + prompt_metadata=prompt_metadata, + mime_type=mime_type, + import_url=True, + ) + + with _results_root(_BLOB_ROOT), _download(content_type="image/png"): + await persist_message_pieces_async(pieces=[piece], serializer_factory=MagicMock(return_value=_serializer())) + + assert piece.prompt_metadata is not None + assert kept.items() <= piece.prompt_metadata.items() + assert piece.prompt_metadata["media_source_url"] == "https://example.test/cat" async def test_converted_url_is_imported_separately_from_text_original() -> None: piece = MessagePieceRequest( - original_value="describe this", converted_value="https://example.test/cat", converted_value_data_type="url" + original_value="describe this", + converted_value="https://example.test/cat", + converted_value_data_type="image_path", + import_url=True, ) factory = MagicMock(return_value=_serializer(value="/stored/cat.png")) @@ -427,18 +483,13 @@ async def test_converted_url_is_imported_separately_from_text_original() -> None assert (piece.data_type, piece.original_value) == ("text", "describe this") assert (piece.converted_value_data_type, piece.converted_value) == ("image_path", "/stored/cat.png") + assert piece.prompt_metadata == { + "converted_media_source_url": "https://example.test/cat", + "converted_media_source_content_type": "image/png", + } -async def test_media_url_piece_keeps_its_declared_type() -> None: - piece = MessagePieceRequest(data_type="audio_path", original_value="https://example.test/clip.wav") - factory = MagicMock(return_value=_serializer(value="/stored/clip.wav")) - - with _results_root(_BLOB_ROOT), _download(content_type="audio/wav"): - await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) - - assert (piece.data_type, piece.original_value, piece.converted_value) == ( - "audio_path", - "/stored/clip.wav", - "/stored/clip.wav", - ) - assert piece.converted_value_data_type is None +@pytest.mark.parametrize("data_type", ["url", "text"]) +def test_import_url_requires_a_declared_media_type(data_type: str) -> None: + with pytest.raises(ValidationError, match="declare the media type"): + MessagePieceRequest(data_type=data_type, original_value="https://example.test/cat", import_url=True) diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py index e9677d84be..7144c1dff2 100644 --- a/tests/unit/backend/test_media_url_import.py +++ b/tests/unit/backend/test_media_url_import.py @@ -3,6 +3,7 @@ """Tests for the bounded media URL download.""" +import logging from collections.abc import Callable, Iterator import httpx @@ -135,24 +136,78 @@ def handler(request: httpx.Request) -> httpx.Response: async def test_error_status_is_rejected_without_query_string( - transport: Callable[[Handler], list[httpx.Request]], + transport: Callable[[Handler], list[httpx.Request]], caplog: pytest.LogCaptureFixture ) -> None: transport(lambda request: httpx.Response(404)) - with pytest.raises(ValueError, match="returned HTTP 404") as error: + with ( + caplog.at_level(logging.WARNING, logger=media_url_import.__name__), + pytest.raises(ValueError, match="returned HTTP 404") as error, + ): await download_media_url_async(url="https://user.example.test/cat.png?sv=1&sig=secret#frag") assert "secret" not in str(error.value) assert "https://user.example.test/cat.png" in str(error.value) + assert "HTTP 404 (HTTPStatusError)" in caplog.text + assert "secret" not in caplog.text + + +_SIGNED_DETAIL = "failed for https://example.test/slow.png?sig=secret" -async def test_network_error_is_rejected(transport: Callable[[Handler], list[httpx.Request]]) -> None: +@pytest.mark.parametrize( + ("make_error", "reason", "cause"), + [ + ( + lambda request: httpx.ConnectTimeout(_SIGNED_DETAIL, request=request), + "connecting timed out after 10 seconds", + "ConnectTimeout", + ), + ( + lambda request: httpx.ReadTimeout(_SIGNED_DETAIL, request=request), + "no data arrived for 30 seconds", + "ReadTimeout", + ), + ( + lambda request: httpx.WriteTimeout(_SIGNED_DETAIL, request=request), + "sending the request timed out after 30 seconds", + "WriteTimeout", + ), + ( + lambda request: httpx.PoolTimeout(_SIGNED_DETAIL, request=request), + "no connection was free within 30 seconds", + "PoolTimeout", + ), + (lambda request: TimeoutError(_SIGNED_DETAIL), "the download took longer than 60 seconds", "TimeoutError"), + (lambda request: httpx.ConnectError(_SIGNED_DETAIL, request=request), "the connection failed", "ConnectError"), + ( + lambda request: httpx.RemoteProtocolError(_SIGNED_DETAIL, request=request), + "the request failed (RemoteProtocolError)", + "RemoteProtocolError", + ), + ], +) +async def test_network_failures_name_the_reason_and_limit( + transport: Callable[[Handler], list[httpx.Request]], + caplog: pytest.LogCaptureFixture, + make_error: Callable[[httpx.Request], BaseException], + reason: str, + cause: str, +) -> None: def fail(request: httpx.Request) -> httpx.Response: - raise httpx.ReadTimeout("slow", request=request) + raise make_error(request) transport(fail) - with pytest.raises(ValueError, match="could not be downloaded"): - await download_media_url_async(url="https://example.test/slow.png") + with ( + caplog.at_level(logging.WARNING, logger=media_url_import.__name__), + pytest.raises(ValueError) as error, + ): + await download_media_url_async(url="https://example.test/slow.png?sig=secret") + + assert str(error.value) == f"Media URL https://example.test/slow.png could not be downloaded: {reason}." + assert error.value.__cause__ is not None + assert f"https://example.test/slow.png could not be downloaded: {reason} ({cause}" in caplog.text + assert "secret" not in caplog.text @pytest.mark.parametrize( diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 8c19d4502b..01eb9a8b72 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -867,13 +867,17 @@ async def test_converted_media_outside_results_is_rejected_async( with pytest.raises(ValueError, match="results directory"): await MessageSendService._persist_base64_pieces_async(request) - async def test_url_piece_is_imported_before_sending_async(self, *, tmp_path: Path) -> None: + async def test_imported_url_is_stored_before_sending_async(self, *, tmp_path: Path) -> None: stored = str(tmp_path / "prompt-memory-entries" / "imported.png") serializer = MagicMock(value=stored) serializer.save_data_async = AsyncMock() download = MediaDownload(content=b"PNG", content_type="image/png", final_url="https://example.com/cat") request = AddMessageRequest( - pieces=[MessagePieceRequest(data_type="url", original_value="https://example.com/cat?sig=secret")], + pieces=[ + MessagePieceRequest( + data_type="image_path", original_value="https://example.com/cat?sig=secret", import_url=True + ) + ], send=True, target_conversation_id="test-id", ) diff --git a/tests/unit/backend/test_runtime_lifecycle.py b/tests/unit/backend/test_runtime_lifecycle.py index 55f4f551d4..a6cf728b64 100644 --- a/tests/unit/backend/test_runtime_lifecycle.py +++ b/tests/unit/backend/test_runtime_lifecycle.py @@ -343,6 +343,19 @@ async def test_management_applies_media_url_import_setting(enabled: bool) -> Non set_enabled.assert_called_once_with(enabled=enabled) +@pytest.mark.parametrize("directory", [None, "/srv/uploads"]) +async def test_management_applies_target_upload_directory(directory: str | None) -> None: + service = RuntimeLifecycle(app=FastAPI(), source=ConfigurationFileService(config_file_value=None)) + config = ConfigurationLoader(memory_db_type="in_memory", env_files=[], target_upload_directory=directory) + with ( + patch.object(lifecycle_module, "set_target_upload_directory") as set_directory, + patch.object(lifecycle_module.InitializerRegistry, "get_registry_singleton", return_value=MagicMock()), + ): + await service._management_async(config) + + set_directory.assert_called_once_with(directory=directory) + + def test_authorization_policy_does_not_change_with_environment(runtime: RuntimeLifecycle) -> None: runtime.app.state.auth_environment["PYRIT_ALLOW_UNAUTHENTICATED_ADMIN"] = "" request = Request({"type": "http", "app": runtime.app}) diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index 4a08dc26b5..8b19c5995e 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -6,10 +6,12 @@ """ import os +from pathlib import Path from unittest.mock import MagicMock, patch import pytest +import pyrit.backend.services.target_service as target_service_module from pyrit.backend.models.targets import CreateTargetRequest from pyrit.backend.services.target_service import TargetService, get_target_service from pyrit.models import ComponentIdentifier @@ -406,26 +408,58 @@ async def test_create_target_raises_for_invalid_type(self) -> None: await service.create_target_async(request=request) @pytest.mark.parametrize( - ("target_type", "params"), + ("target_type", "params", "upload_directory", "error"), [ - ("HTTPXAPITarget", {"http_url": "http://localhost:8080/upload"}), - ("HuggingFaceChatTarget", {"model_id": "example/model"}), + ("HuggingFaceChatTarget", {"model_id": "example/model"}, None, "loads model code"), + ("HTTPXAPITarget", {"http_url": "http://localhost:8080/upload"}, None, "target_upload_directory"), + ( + "HTTPXAPITarget", + {"http_url": "http://localhost:8080/upload", "allowed_upload_directory": "/"}, + "configured", + "'allowed_upload_directory' of 'HTTPXAPITarget' names a path on this server", + ), + ( + "GitHubCopilotTarget", + {"model_name": "gpt-5", "working_directory": "/"}, + None, + "'working_directory' of 'GitHubCopilotTarget' names a path on this server", + ), ], ) - async def test_create_target_rejects_types_using_host_resources( - self, sqlite_instance, target_type: str, params: dict[str, str] + async def test_create_target_rejects_server_resources( + self, + sqlite_instance, + tmp_path: Path, + target_type: str, + params: dict[str, str], + upload_directory: str | None, + error: str, ) -> None: service = TargetService() request = CreateTargetRequest(name="host-target", type=target_type, params=params) + configured_directory = str(tmp_path) if upload_directory else None with ( - patch.object(service._registry, "create_named_instance") as create, - pytest.raises(ValueError, match="cannot be created through the API"), + patch.object(target_service_module, "_target_upload_directory", configured_directory), + patch.object(service._registry, "create_instance") as create, + pytest.raises(ValueError, match=error), ): await service.create_target_async(request=request) create.assert_not_called() + async def test_create_upload_target_uses_the_configured_directory(self, sqlite_instance, tmp_path: Path) -> None: + service = TargetService() + request = CreateTargetRequest( + name="uploader", type="HTTPXAPITarget", params={"http_url": "http://localhost:8080/upload"} + ) + + with patch.object(target_service_module, "_target_upload_directory", str(tmp_path)): + await service.create_target_async(request=request) + + target = service.get_target_object(target_registry_name="uploader") + assert target.allowed_upload_directory == tmp_path.resolve() + async def test_create_target_success(self, sqlite_instance) -> None: """Test successful target creation.""" service = TargetService() diff --git a/tests/unit/setup/test_configuration_loader.py b/tests/unit/setup/test_configuration_loader.py index 7bee2b160c..6f3ace97c9 100644 --- a/tests/unit/setup/test_configuration_loader.py +++ b/tests/unit/setup/test_configuration_loader.py @@ -85,6 +85,16 @@ def test_rejects_non_boolean_allow_media_url_import(self, invalid_value: object) with pytest.raises(TypeError, match=r"allow_media_url_import must be a bool"): ConfigurationLoader(allow_media_url_import=invalid_value) # type: ignore[arg-type] + def test_target_upload_directory_is_unset_by_default(self) -> None: + assert ConfigurationLoader().target_upload_directory is None + config = ConfigurationLoader.from_dict({"target_upload_directory": "/srv/uploads"}) + assert config.target_upload_directory == "/srv/uploads" + + @pytest.mark.parametrize("invalid_value", [1, True, [], {}]) + def test_rejects_non_string_target_upload_directory(self, invalid_value: object) -> None: + with pytest.raises(TypeError, match=r"target_upload_directory must be a string path"): + ConfigurationLoader(target_upload_directory=invalid_value) # type: ignore[arg-type] + def test_valid_memory_db_types_snake_case(self): """Test all valid memory database types in snake_case.""" for db_type in ["in_memory", "sqlite", "azure_sql"]: From c4a7431b6d10242478e08dd58e55506556a4613f Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 19:18:39 +0000 Subject: [PATCH 4/7] Keep URL queries out of import tracebacks and list only targets the API can create Import failures are raised without the httpx exception chain, since httpx messages quote the full URL and background sends log the traceback. A URL httpx cannot parse is a 400 instead of a 500. A blank target_upload_directory is rejected rather than resolving to the working directory, and the target type catalog leaves out types and path parameters that creation always rejects. --- pyrit/backend/README.md | 10 ++-- pyrit/backend/services/media_url_import.py | 8 +++- pyrit/backend/services/target_service.py | 47 ++++++++++++++----- pyrit/setup/configuration_loader.py | 2 + tests/unit/backend/test_media_url_import.py | 6 ++- .../unit/backend/test_message_send_service.py | 38 +++++++++++++++ tests/unit/backend/test_target_service.py | 25 ++++++++-- tests/unit/setup/test_configuration_loader.py | 5 ++ 8 files changed, 116 insertions(+), 25 deletions(-) diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index eca59a81d8..e057243b10 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -225,11 +225,11 @@ request values before using them: in `.pyrit_conf` to turn imports off. Converter file parameters given a URL are downloaded the same way. - Target types that load model code (`HuggingFaceChatTarget`) cannot be created through the - API, and target parameters that name server paths cannot be set through it. Targets that - upload local files (`HTTPXAPITarget`) can be created through the API only when - `target_upload_directory` is set in `.pyrit_conf`; the server passes that directory to the - target, which uploads files only from inside it. Register such targets in Python or with an - initializer for other settings. + API, and target parameters that name server paths cannot be set through it; the target type + catalog leaves both out. Targets that upload local files (`HTTPXAPITarget`) can be created + through the API only when `target_upload_directory` is set in `.pyrit_conf`; the server passes + that directory to the target, which uploads files only from inside it. Register such targets + in Python or with an initializer for other settings. Intentional exceptions: diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py index 8d6a7b7e1b..3a6f0e5a94 100644 --- a/pyrit/backend/services/media_url_import.py +++ b/pyrit/backend/services/media_url_import.py @@ -171,13 +171,17 @@ async def download_media_url_async(*, url: str) -> MediaDownload: if not _is_plain_http_url(request.url): raise ValueError(f"Media URL {shown} redirected to a URL that is not a plain http or https URL.") raise ValueError(f"Media URL {shown} redirected more than {MAX_MEDIA_URL_REDIRECTS} times.") + # The httpx exceptions are not chained because their messages quote the full URL, query string included, + # and callers may log the raised error with its traceback. The failure is logged here with the redacted URL. + except httpx.InvalidURL: + raise ValueError(f"Media URL {shown} is not a valid http or https URL.") from None except httpx.HTTPStatusError as exc: _log_download_failure(shown=shown, reason=f"HTTP {exc.response.status_code}", exc=exc) - raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from exc + raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from None except (httpx.HTTPError, TimeoutError) as exc: reason = _failure_reason(exc) _log_download_failure(shown=shown, reason=reason, exc=exc) - raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from exc + raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from None def _failure_reason(exc: BaseException) -> str: diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 1d7a457d9e..f09550edbf 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -53,6 +53,22 @@ def set_target_upload_directory(*, directory: str | None) -> None: _target_upload_directory = directory +def _names_server_path(parameter: Parameter) -> bool: + return parameter.is_path or parameter.is_path_or_str + + +def _can_create_through_api(target_cls: type[PromptTarget]) -> bool: + """ + Return whether the API may create a target type under the current server configuration. + + Returns: + bool: False for types that load model code, or upload local files while no upload directory is configured. + """ + if target_cls.loads_local_code: + return False + return not (target_cls.upload_directory_parameter and _target_upload_directory is None) + + class TargetService: """ Service for managing target instances. @@ -221,7 +237,7 @@ def _reject_server_resources( "with an initializer instead." ) for parameter in parameters: - if parameter.name in params and (parameter.is_path or parameter.is_path_or_str): + if parameter.name in params and _names_server_path(parameter): raise ValueError( f"Parameter '{parameter.name}' of '{target_type}' names a path on this server and cannot be " "set through the API." @@ -231,14 +247,14 @@ async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. - Returns every registered target with its derived constructor - parameters and the auth modes it supports, all projected from the - registry's ``TargetMetadata``. Types that load model code, or upload - local files when the server has no upload directory configured, are - listed but cannot be created through the API, and neither can parameters - that name server paths. Deciding which entries to surface to a user is a - presentation concern owned by the caller (e.g. the frontend), not this - service. + Returns every target the API can create under the server configuration, + with the constructor parameters callers may supply and the auth modes it + supports, all projected from the registry's ``TargetMetadata``. Types that + load model code, or upload local files when the server has no upload + directory configured, are left out, and so are parameters that name server + paths; the registry metadata keeps them. Deciding which entries to surface to + a user is a presentation concern owned by the caller (e.g. the frontend), not + this service. Returns: TargetTypeResponse containing all available target classes. @@ -247,14 +263,19 @@ async def list_target_types_async(self) -> TargetTypeResponse: items: list[TargetTypeEntry] = [ TargetTypeEntry( target_type=metadata.class_name, - parameters=self._project_target_parameters( - target_type=metadata.class_name, - parameters=metadata.parameters, - ), + parameters=[ + parameter + for parameter in self._project_target_parameters( + target_type=metadata.class_name, + parameters=metadata.parameters, + ) + if not _names_server_path(parameter) + ], supported_auth_modes=self._get_supported_auth_modes(metadata.supported_auth_modes), description=metadata.class_description or None, ) for metadata in metadata_items + if _can_create_through_api(self._registry.get_class(metadata.class_name)) ] return TargetTypeResponse(items=items) diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index bf6ab9dd90..c2559bc086 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -195,6 +195,8 @@ def __post_init__(self) -> None: raise TypeError("allow_media_url_import must be a bool.") if self.target_upload_directory is not None and not isinstance(self.target_upload_directory, str): raise TypeError("target_upload_directory must be a string path.") + if self.target_upload_directory is not None and not is_non_empty_string(self.target_upload_directory): + raise ValueError("target_upload_directory must be a non-empty path.") self._validate_allow_custom_initializers() self._normalize_memory_db_type() self._normalize_initializers() diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py index 7144c1dff2..f3bedcaa11 100644 --- a/tests/unit/backend/test_media_url_import.py +++ b/tests/unit/backend/test_media_url_import.py @@ -146,6 +146,8 @@ async def test_error_status_is_rejected_without_query_string( ): await download_media_url_async(url="https://user.example.test/cat.png?sv=1&sig=secret#frag") assert "secret" not in str(error.value) + assert error.value.__cause__ is None + assert error.value.__suppress_context__ assert "https://user.example.test/cat.png" in str(error.value) assert "HTTP 404 (HTTPStatusError)" in caplog.text assert "secret" not in caplog.text @@ -205,7 +207,8 @@ def fail(request: httpx.Request) -> httpx.Response: await download_media_url_async(url="https://example.test/slow.png?sig=secret") assert str(error.value) == f"Media URL https://example.test/slow.png could not be downloaded: {reason}." - assert error.value.__cause__ is not None + assert error.value.__cause__ is None + assert error.value.__suppress_context__ assert f"https://example.test/slow.png could not be downloaded: {reason} ({cause}" in caplog.text assert "secret" not in caplog.text @@ -215,6 +218,7 @@ def fail(request: httpx.Request) -> httpx.Response: [ ("ftp://example.test/cat.png", "http or https"), ("https:///cat.png", "http or https"), + ("https://exa\x01mple.test/cat.png?sig=secret", "not a valid http or https URL"), ("https://user:pass@example.test/cat.png", "credentials"), ], ) diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 01eb9a8b72..90b9fa0ceb 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -4,6 +4,7 @@ """Tests for the shared synchronous manual-message owner.""" import asyncio +import traceback import uuid from collections.abc import AsyncGenerator, Generator, Iterator, Sequence from contextlib import asynccontextmanager, contextmanager @@ -12,6 +13,7 @@ from typing import Any from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from pydantic import ValidationError from sqlalchemy.orm import Session @@ -3208,6 +3210,42 @@ async def test_preparation_failure_uses_dispatch_stage_not_saved_messages_async( assert status.failure_stage == MessageSendFailureStage.PREPARATION assert status.error + @pytest.mark.parametrize("count", [1, 2]) + async def test_failed_url_import_logs_no_signed_query_async( + self, + *, + real_send_context: tuple[MessageSendService, AttackResult, MockPromptTarget, Base64Converter], + caplog: pytest.LogCaptureFixture, + count: int, + ) -> None: + service, ar, _, _ = real_send_context + request = MessageSendRequest( + pieces=[ + MessagePieceRequest( + data_type="image_path", original_value="https://example.test/cat.png?sig=secret", import_url=True + ) + ], + target_conversation_id=ar.conversation_id, + target_registry_name="target", + send=True, + submission_id="submission", + count=count, + ) + + def create_client() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(lambda _: httpx.Response(404))) + + with patch("pyrit.backend.services.media_url_import._create_client", create_client): + status = await service.submit_async(attack_result_id=ar.attack_result_id, request=request) + status = await _settle_send_async(service=service, status=status) + + assert status.state == MessageSendState.FAILED + assert status.failure_stage == MessageSendFailureStage.PREPARATION + [failure] = [record for record in caplog.records if record.exc_info] + assert "returned HTTP 404" in str(failure.exc_info[1]) + assert "secret" not in "".join(traceback.format_exception(*failure.exc_info)) + assert "secret" not in caplog.text + @pytest.mark.parametrize("count", [1, 3]) @pytest.mark.parametrize("failure", ["conversion", "normalization", "validation", "target", "metadata"]) async def test_real_sqlite_failure_stages_preserve_evidence_async( diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index 8b19c5995e..f0a18be66d 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -265,7 +265,6 @@ async def test_types_include_declarative_auth_facts(self) -> None: ("OpenAIChatTarget", {"endpoint", "model_name"}), ("AzureBlobStorageTarget", {"container_url"}), ("HackAPromptTarget", {"cookie", "session_id"}), - ("HuggingFaceChatTarget", {"hf_access_token"}), ("PromptShieldTarget", {"endpoint"}), ("AzureMLChatTarget", {"endpoint"}), ], @@ -301,21 +300,39 @@ async def test_types_include_structured_parameters(self) -> None: async def test_types_preserve_registry_parameter_order_without_mutating_metadata(self) -> None: service = TargetService() - result = await service.list_target_types_async() + with patch.object(target_service_module, "_target_upload_directory", None): + 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) + assert {entry.target_type for entry in result.items} == set(metadata_by_name) - { + "HuggingFaceChatTarget", + "HTTPXAPITarget", + } for entry in result.items: registry_parameters = metadata_by_name[entry.target_type].parameters assert [parameter.name for parameter in entry.parameters] == [ - parameter.name for parameter in registry_parameters + parameter.name for parameter in registry_parameters if parameter.name != "working_directory" ] registry_openai = {parameter.name: parameter for parameter in metadata_by_name["OpenAIChatTarget"].parameters} assert registry_openai["endpoint"].required is False assert registry_openai["model_name"].required is False + registry_copilot = {parameter.name for parameter in metadata_by_name["GitHubCopilotTarget"].parameters} + assert "working_directory" in registry_copilot + + async def test_types_list_the_upload_target_without_its_directory_once_configured(self, tmp_path: Path) -> None: + service = TargetService() + + with patch.object(target_service_module, "_target_upload_directory", str(tmp_path)): + result = await service.list_target_types_async() + + entry = next(item for item in result.items if item.target_type == "HTTPXAPITarget") + names = [parameter.name for parameter in entry.parameters] + assert "file_path" in names + assert "allowed_upload_directory" not in names + assert "HuggingFaceChatTarget" not in {item.target_type for item in result.items} async def test_types_cold_and_warm_results_are_equal(self) -> None: service = TargetService() diff --git a/tests/unit/setup/test_configuration_loader.py b/tests/unit/setup/test_configuration_loader.py index 6f3ace97c9..a91eb700a8 100644 --- a/tests/unit/setup/test_configuration_loader.py +++ b/tests/unit/setup/test_configuration_loader.py @@ -95,6 +95,11 @@ def test_rejects_non_string_target_upload_directory(self, invalid_value: object) with pytest.raises(TypeError, match=r"target_upload_directory must be a string path"): ConfigurationLoader(target_upload_directory=invalid_value) # type: ignore[arg-type] + @pytest.mark.parametrize("blank_value", ["", " "]) + def test_rejects_blank_target_upload_directory(self, blank_value: str) -> None: + with pytest.raises(ValueError, match=r"target_upload_directory must be a non-empty path"): + ConfigurationLoader(target_upload_directory=blank_value) + def test_valid_memory_db_types_snake_case(self): """Test all valid memory database types in snake_case.""" for db_type in ["in_memory", "sqlite", "azure_sql"]: From 8371dc8e2992400385679385c7afed1f9ec494cb Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 20:12:24 +0000 Subject: [PATCH 5/7] Redact media URLs in httpx request logs during import httpx logs every request URL at INFO, query string included. While a media URL downloads, a filter on the httpx logger replaces those URLs with the redacted form; other httpx logging is unchanged. --- pyrit/backend/services/media_url_import.py | 17 +++++++++++++++++ tests/unit/backend/test_media_url_import.py | 17 +++++++++++++++++ tests/unit/backend/test_message_send_service.py | 2 ++ 3 files changed, 36 insertions(+) diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py index 3a6f0e5a94..727b1f1bc9 100644 --- a/pyrit/backend/services/media_url_import.py +++ b/pyrit/backend/services/media_url_import.py @@ -9,6 +9,7 @@ import logging import mimetypes import re +from contextvars import ContextVar from dataclasses import dataclass from pathlib import PurePosixPath from urllib.parse import urlparse @@ -29,6 +30,7 @@ _URL_SUFFIX_PATTERN = re.compile(r"^\.[A-Za-z0-9]{1,10}$") _url_import_enabled = True +_redact_request_log: ContextVar[bool] = ContextVar("redact_media_url_request_log", default=False) @dataclass(frozen=True) @@ -60,6 +62,18 @@ def redact_url(url: str) -> str: return f"{parsed.scheme}://{host}{parsed.path}" +class _RedactRequestLog(logging.Filter): + """Redact the URLs httpx logs for each request while a media URL is downloading.""" + + def filter(self, record: logging.LogRecord) -> bool: + if _redact_request_log.get() and isinstance(record.args, tuple): + record.args = tuple(redact_url(str(arg)) if isinstance(arg, httpx.URL) else arg for arg in record.args) + return True + + +logging.getLogger("httpx").addFilter(_RedactRequestLog()) + + def media_content_type(download: MediaDownload) -> str | None: """ Return the reported media type, or the type implied by the URL suffix when it is missing or generic. @@ -155,6 +169,7 @@ async def download_media_url_async(*, url: str) -> MediaDownload: raise ValueError("Media URLs must not include credentials.") shown = redact_url(url) + redacting_request_log = _redact_request_log.set(True) try: async with asyncio.timeout(_DEADLINE_SECONDS), _create_client() as client: request = client.build_request("GET", url) @@ -182,6 +197,8 @@ async def download_media_url_async(*, url: str) -> MediaDownload: reason = _failure_reason(exc) _log_download_failure(shown=shown, reason=reason, exc=exc) raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from None + finally: + _redact_request_log.reset(redacting_request_log) def _failure_reason(exc: BaseException) -> str: diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py index f3bedcaa11..3d2010f435 100644 --- a/tests/unit/backend/test_media_url_import.py +++ b/tests/unit/backend/test_media_url_import.py @@ -142,6 +142,7 @@ async def test_error_status_is_rejected_without_query_string( with ( caplog.at_level(logging.WARNING, logger=media_url_import.__name__), + caplog.at_level(logging.INFO, logger="httpx"), pytest.raises(ValueError, match="returned HTTP 404") as error, ): await download_media_url_async(url="https://user.example.test/cat.png?sv=1&sig=secret#frag") @@ -150,9 +151,25 @@ async def test_error_status_is_rejected_without_query_string( assert error.value.__suppress_context__ assert "https://user.example.test/cat.png" in str(error.value) assert "HTTP 404 (HTTPStatusError)" in caplog.text + assert "HTTP Request: GET https://user.example.test/cat.png " in caplog.text assert "secret" not in caplog.text +async def test_request_log_omits_query_string_only_while_downloading( + transport: Callable[[Handler], list[httpx.Request]], caplog: pytest.LogCaptureFixture +) -> None: + transport(lambda request: httpx.Response(200, content=b"PNG", headers={"content-type": "image/png"})) + + with caplog.at_level(logging.INFO, logger="httpx"): + await download_media_url_async(url="https://example.test/cat.png?sig=secret") + async with httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(200))) as client: + await client.get("https://other.test/page?keep=1") + + assert "HTTP Request: GET https://example.test/cat.png " in caplog.text + assert "secret" not in caplog.text + assert "HTTP Request: GET https://other.test/page?keep=1 " in caplog.text + + _SIGNED_DETAIL = "failed for https://example.test/slow.png?sig=secret" diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 90b9fa0ceb..34b8ca8808 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -4,6 +4,7 @@ """Tests for the shared synchronous manual-message owner.""" import asyncio +import logging import traceback import uuid from collections.abc import AsyncGenerator, Generator, Iterator, Sequence @@ -3219,6 +3220,7 @@ async def test_failed_url_import_logs_no_signed_query_async( count: int, ) -> None: service, ar, _, _ = real_send_context + caplog.set_level(logging.INFO) request = MessageSendRequest( pieces=[ MessagePieceRequest( From 805d843a93560bfc6c76afbe56ae8ae675258650 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 8 Oct 2026 20:37:16 +0000 Subject: [PATCH 6/7] Drop httpcore protocol traces while a media URL downloads httpcore logs response headers at DEBUG, so a signed redirect Location or Content-Location would reach the log. Those traces are dropped during imports; other httpcore logging is unchanged. --- pyrit/backend/services/media_url_import.py | 18 +++++++++++++++--- tests/unit/backend/test_media_url_import.py | 12 ++++++++++-- 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py index 727b1f1bc9..ddc1995a0d 100644 --- a/pyrit/backend/services/media_url_import.py +++ b/pyrit/backend/services/media_url_import.py @@ -63,15 +63,27 @@ def redact_url(url: str) -> str: class _RedactRequestLog(logging.Filter): - """Redact the URLs httpx logs for each request while a media URL is downloading.""" + """ + Keep media URL query strings out of HTTP client logs while a media URL is downloading. + + httpx request lines are logged with the redacted URL. httpcore's protocol traces are dropped, + because they include response headers such as a signed redirect ``Location``. + """ def filter(self, record: logging.LogRecord) -> bool: - if _redact_request_log.get() and isinstance(record.args, tuple): + if not _redact_request_log.get(): + return True + if record.name.startswith("httpcore."): + return False + if isinstance(record.args, tuple): record.args = tuple(redact_url(str(arg)) if isinstance(arg, httpx.URL) else arg for arg in record.args) return True -logging.getLogger("httpx").addFilter(_RedactRequestLog()) +_request_log_filter = _RedactRequestLog() +logging.getLogger("httpx").addFilter(_request_log_filter) +logging.getLogger("httpcore.http11").addFilter(_request_log_filter) +logging.getLogger("httpcore.http2").addFilter(_request_log_filter) def media_content_type(download: MediaDownload) -> str | None: diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py index 3d2010f435..5c2b7c5379 100644 --- a/tests/unit/backend/test_media_url_import.py +++ b/tests/unit/backend/test_media_url_import.py @@ -158,16 +158,24 @@ async def test_error_status_is_rejected_without_query_string( async def test_request_log_omits_query_string_only_while_downloading( transport: Callable[[Handler], list[httpx.Request]], caplog: pytest.LogCaptureFixture ) -> None: - transport(lambda request: httpx.Response(200, content=b"PNG", headers={"content-type": "image/png"})) + trace_logger = logging.getLogger("httpcore.http11") - with caplog.at_level(logging.INFO, logger="httpx"): + def respond(request: httpx.Request) -> httpx.Response: + trace_logger.debug("receive_response_headers.complete Location=https://example.test/next.png?sig=secret") + return httpx.Response(200, content=b"PNG", headers={"content-type": "image/png"}) + + transport(respond) + + with caplog.at_level(logging.DEBUG, logger="httpx"), caplog.at_level(logging.DEBUG, logger="httpcore"): await download_media_url_async(url="https://example.test/cat.png?sig=secret") + trace_logger.debug("receive_response_headers.complete Location=https://other.test/next?keep=1") async with httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(200))) as client: await client.get("https://other.test/page?keep=1") assert "HTTP Request: GET https://example.test/cat.png " in caplog.text assert "secret" not in caplog.text assert "HTTP Request: GET https://other.test/page?keep=1 " in caplog.text + assert "Location=https://other.test/next?keep=1" in caplog.text _SIGNED_DETAIL = "failed for https://example.test/slow.png?sig=secret" From 2d4aa44b8e9c973c97798b7caa26091f2e264fbe Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 8 Oct 2026 17:27:36 -0700 Subject: [PATCH 7/7] Fix imported media formats and restore HTTP diagnostics Resolve imported formats from response and caller MIME types or URL suffixes without changing bytes or defaulting unknown media to WAV. Remove global HTTP logging filters and preserve download exception causes. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/backend/README.md | 7 +- pyrit/backend/services/media_persistence.py | 35 +++----- pyrit/backend/services/media_url_import.py | 90 +++++++------------ tests/unit/backend/test_media_persistence.py | 82 ++++++++++++++++- tests/unit/backend/test_media_url_import.py | 30 +++---- .../unit/backend/test_message_send_service.py | 6 +- 6 files changed, 139 insertions(+), 111 deletions(-) diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index e057243b10..9994e3e17d 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -218,9 +218,12 @@ request values before using them: preview with an `image_path`, `audio_path`, `video_path`, or `binary_path` type. The server downloads it once into managed storage (10 second connect, 30 second read, and 60 second total limits, 100 MiB limit, at most 3 redirects, no request credentials forwarded) and - stores it under the declared type. Converters and targets then see only the stored copy, + stores it under the declared type without format conversion. The format extension comes + from the response MIME type, the caller MIME type if the response is missing or generic, + or the URL suffix. Unknown formats use `.bin`, not a modality default such as `.wav`. + Converters and targets then see only the stored copy, and the piece's prompt metadata records the source URL, without credentials or query - string, and the reported content type. A preview that imports returns the stored copy and + string, and the resolved content type. A preview that imports returns the stored copy and that metadata, so sending them reuses the same bytes. Set `allow_media_url_import: false` in `.pyrit_conf` to turn imports off. Converter file parameters given a URL are downloaded the same way. diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index 29c2f35fd6..41349fa204 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -18,9 +18,9 @@ from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS from pyrit.backend.services.media_url_import import ( - MediaDownload, download_media_url_async, media_content_type, + media_extension, redact_url, ) from pyrit.common.azure_storage import is_azure_blob_uri, redact_url_credentials @@ -254,27 +254,11 @@ async def _write_owned_media_async( raise -def _imported_extension(*, download: MediaDownload, data_type: PromptDataType) -> str: - """ - Choose the file extension for imported media of a declared path type. - - Returns: - str: The extension of the reported (or URL-implied) media type when it fits the declared - type, else the declared type's default extension. - """ - content_type = media_content_type(download) - expected_family = _MEDIA_FAMILIES.get(data_type) - if content_type and expected_family in (None, content_type.split("/", 1)[0]): - extension = mimetypes.guess_extension(content_type, strict=False) - if extension: - return extension - return DEFAULT_MEDIA_EXTENSIONS.get(str(data_type), ".bin") - - async def _import_media_url_async( *, url: str, data_type: PromptDataType, + mime_type: str | None, serializer_factory: SerializerFactory, created_paths: list[str] | None, ) -> MediaPersistenceResult: @@ -282,7 +266,8 @@ async def _import_media_url_async( Download a media URL once and store the bytes in managed media storage under the declared type. A missing or generic reported content type is accepted, since the caller declared the type; - content of a different media family than the declared type is rejected. + content of a different media family than the declared type is rejected. Format is resolved + separately from the response, caller MIME type, or URL suffix; unknown formats use ``.bin``. Returns: MediaPersistenceResult: The managed reference to the stored copy and its redacted source URL. @@ -291,12 +276,12 @@ async def _import_media_url_async( ValueError: If the download fails or the content does not match the declared media type. """ download = await download_media_url_async(url=url) - content_type = media_content_type(download) + content_type = media_content_type(download, mime_type=mime_type) expected_family = _MEDIA_FAMILIES.get(data_type) family = (content_type or "").split("/", 1)[0] if expected_family and family in _CHECKED_FAMILIES and family != expected_family: raise ValueError(f"Media URL {redact_url(url)} returned {content_type}, not {expected_family} content.") - extension = _imported_extension(download=download, data_type=data_type) + extension = media_extension(download, mime_type=mime_type, default=".bin") serializer = serializer_factory( category="prompt-memory-entries", data_type=data_type, @@ -375,7 +360,11 @@ async def persist_media_value_async( ) if import_url: return await _import_media_url_async( - url=value, data_type=data_type, serializer_factory=serializer_factory, created_paths=created_paths + url=value, + data_type=data_type, + mime_type=mime_type, + serializer_factory=serializer_factory, + created_paths=created_paths, ) if _is_read_from_result_storage(value): raise ValueError( @@ -469,7 +458,7 @@ def media_source_metadata(result: MediaPersistenceResult, *, prefix: str = "medi prefix (str): Key prefix, so a piece can record its original and converted sources apart. Returns: - dict[str, str]: The redacted source URL and the content type the server reported, or an + dict[str, str]: The redacted source URL and the resolved content type, or an empty dict when the value was not imported. """ if result.source is None: diff --git a/pyrit/backend/services/media_url_import.py b/pyrit/backend/services/media_url_import.py index ddc1995a0d..4f47342132 100644 --- a/pyrit/backend/services/media_url_import.py +++ b/pyrit/backend/services/media_url_import.py @@ -6,10 +6,8 @@ from __future__ import annotations import asyncio -import logging import mimetypes import re -from contextvars import ContextVar from dataclasses import dataclass from pathlib import PurePosixPath from urllib.parse import urlparse @@ -18,8 +16,6 @@ from pyrit.common.net_utility import get_httpx_client -logger = logging.getLogger(__name__) - MAX_MEDIA_URL_BYTES = 100 * 1024 * 1024 MAX_MEDIA_URL_REDIRECTS = 3 _CONNECT_TIMEOUT_SECONDS = 10.0 @@ -28,9 +24,19 @@ _DEADLINE_SECONDS = 60.0 _GENERIC_CONTENT_TYPES = frozenset({"application/octet-stream", "binary/octet-stream"}) _URL_SUFFIX_PATTERN = re.compile(r"^\.[A-Za-z0-9]{1,10}$") +_MEDIA_EXTENSIONS = { + "audio/mpeg": ".mp3", + "audio/mp3": ".mp3", + "audio/wav": ".wav", + "audio/x-wav": ".wav", + "audio/flac": ".flac", + "audio/x-flac": ".flac", + "audio/ogg": ".ogg", + "video/ogg": ".ogv", + "application/ogg": ".ogg", +} _url_import_enabled = True -_redact_request_log: ContextVar[bool] = ContextVar("redact_media_url_request_log", default=False) @dataclass(frozen=True) @@ -62,53 +68,36 @@ def redact_url(url: str) -> str: return f"{parsed.scheme}://{host}{parsed.path}" -class _RedactRequestLog(logging.Filter): - """ - Keep media URL query strings out of HTTP client logs while a media URL is downloading. - - httpx request lines are logged with the redacted URL. httpcore's protocol traces are dropped, - because they include response headers such as a signed redirect ``Location``. +def media_content_type(download: MediaDownload, *, mime_type: str | None = None) -> str | None: """ + Resolve the media type from the response, caller MIME type, then URL suffix. - def filter(self, record: logging.LogRecord) -> bool: - if not _redact_request_log.get(): - return True - if record.name.startswith("httpcore."): - return False - if isinstance(record.args, tuple): - record.args = tuple(redact_url(str(arg)) if isinstance(arg, httpx.URL) else arg for arg in record.args) - return True - - -_request_log_filter = _RedactRequestLog() -logging.getLogger("httpx").addFilter(_request_log_filter) -logging.getLogger("httpcore.http11").addFilter(_request_log_filter) -logging.getLogger("httpcore.http2").addFilter(_request_log_filter) - - -def media_content_type(download: MediaDownload) -> str | None: - """ - Return the reported media type, or the type implied by the URL suffix when it is missing or generic. + Missing or generic types do not identify an encoding and are skipped. Returns: str | None: The lowercase media type without parameters, or None when unknown. """ - content_type = (download.content_type or "").split(";", 1)[0].strip().lower() - if content_type and content_type not in _GENERIC_CONTENT_TYPES: - return content_type + for candidate in (download.content_type, mime_type): + content_type = (candidate or "").split(";", 1)[0].strip().lower() + if content_type and content_type not in _GENERIC_CONTENT_TYPES: + return content_type guessed, _ = mimetypes.guess_type(urlparse(download.final_url).path, strict=False) return guessed -def media_extension(download: MediaDownload, *, default: str) -> str: +def media_extension(download: MediaDownload, *, default: str, mime_type: str | None = None) -> str: """ Choose the file extension for downloaded media. Returns: str: The extension implied by the media type, else a short suffix from the URL path, else ``default``. """ - content_type = media_content_type(download) - extension = mimetypes.guess_extension(content_type, strict=False) if content_type else None + content_type = media_content_type(download, mime_type=mime_type) + extension = ( + _MEDIA_EXTENSIONS.get(content_type) or mimetypes.guess_extension(content_type, strict=False) + if content_type + else None + ) if extension: return extension suffix = PurePosixPath(urlparse(download.final_url).path).suffix @@ -181,7 +170,6 @@ async def download_media_url_async(*, url: str) -> MediaDownload: raise ValueError("Media URLs must not include credentials.") shown = redact_url(url) - redacting_request_log = _redact_request_log.set(True) try: async with asyncio.timeout(_DEADLINE_SECONDS), _create_client() as client: request = client.build_request("GET", url) @@ -198,26 +186,18 @@ async def download_media_url_async(*, url: str) -> MediaDownload: if not _is_plain_http_url(request.url): raise ValueError(f"Media URL {shown} redirected to a URL that is not a plain http or https URL.") raise ValueError(f"Media URL {shown} redirected more than {MAX_MEDIA_URL_REDIRECTS} times.") - # The httpx exceptions are not chained because their messages quote the full URL, query string included, - # and callers may log the raised error with its traceback. The failure is logged here with the redacted URL. - except httpx.InvalidURL: - raise ValueError(f"Media URL {shown} is not a valid http or https URL.") from None + except httpx.InvalidURL as exc: + raise ValueError(f"Media URL {shown} is not a valid http or https URL.") from exc except httpx.HTTPStatusError as exc: - _log_download_failure(shown=shown, reason=f"HTTP {exc.response.status_code}", exc=exc) - raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from None + raise ValueError(f"Media URL {shown} returned HTTP {exc.response.status_code}.") from exc except (httpx.HTTPError, TimeoutError) as exc: reason = _failure_reason(exc) - _log_download_failure(shown=shown, reason=reason, exc=exc) - raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from None - finally: - _redact_request_log.reset(redacting_request_log) + raise ValueError(f"Media URL {shown} could not be downloaded: {reason}.") from exc def _failure_reason(exc: BaseException) -> str: """ - Describe why a download failed, with the limit that was hit, without quoting the exception. - - httpx exception messages can include the full URL and its query string, so they are not repeated. + Describe why a download failed, with the limit that was hit. Returns: str: A short reason such as ``connecting timed out after 10 seconds``. @@ -235,13 +215,3 @@ def _failure_reason(exc: BaseException) -> str: if isinstance(exc, httpx.ConnectError): return "the connection failed" return f"the request failed ({type(exc).__name__})" - - -def _log_download_failure(*, shown: str, reason: str, exc: BaseException) -> None: - """Log a failed download with the redacted URL and the exception classes in its cause chain.""" - causes: list[str] = [] - cause: BaseException | None = exc - while cause is not None and len(causes) < 4: - causes.append(type(cause).__name__) - cause = cause.__cause__ or cause.__context__ - logger.warning("Media URL %s could not be downloaded: %s (%s)", shown, reason, " <- ".join(causes)) diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index 125162badc..19f680b4a8 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -3,10 +3,13 @@ """Tests for shared backend media persistence.""" +import base64 from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch from urllib.parse import quote +import aiofiles +import httpx import pytest from pydantic import ValidationError @@ -19,8 +22,11 @@ require_managed_blob_url, ) from pyrit.backend.services.media_url_import import MediaDownload -from pyrit.memory import CentralMemory +from pyrit.memory import CentralMemory, SQLiteMemory from pyrit.memory.storage.storage import AzureBlobStorageIO +from pyrit.models import Message, MessagePiece +from pyrit.prompt_target import HTTPXAPITarget +from pyrit.prompt_target.common.chat_completions_message_builder import build_audio_content_entry_async _BLOB_ROOT = "https://account.blob.core.windows.net/results" @@ -182,10 +188,10 @@ async def test_import_of_managed_blob_url_keeps_the_reference() -> None: @pytest.mark.parametrize( ("data_type", "content_type", "final_url", "expected_extension"), [ - ("image_path", "application/octet-stream", "https://example.test/a", ".png"), - ("image_path", None, "https://example.test/a", ".png"), + ("image_path", "application/octet-stream", "https://example.test/a", ".bin"), + ("image_path", None, "https://example.test/a", ".bin"), ("image_path", "application/octet-stream", "https://example.test/a/photo.jpg", ".jpg"), - ("audio_path", "application/ogg", "https://example.test/a", ".wav"), + ("audio_path", "application/ogg", "https://example.test/a", ".ogg"), ("binary_path", "application/pdf", "https://example.test/a", ".pdf"), ("binary_path", "image/png", "https://example.test/a", ".png"), ], @@ -203,6 +209,74 @@ async def test_import_stores_the_declared_type( factory.assert_called_once_with(category="prompt-memory-entries", data_type=data_type, extension=expected_extension) +@pytest.mark.parametrize( + ("content_type", "mime_type", "final_url", "content", "extension", "upload_type", "chat_format"), + [ + ("application/octet-stream", "audio/mpeg", "https://example.test/a", b"ID3-MP3", ".mp3", "audio/mpeg", "mp3"), + (None, " AUDIO/MPEG; charset=binary ", "https://example.test/a", b"ID3-MP3", ".mp3", "audio/mpeg", "mp3"), + ("audio/mpeg", "audio/wav", "https://example.test/a", b"ID3-MP3", ".mp3", "audio/mpeg", "mp3"), + ("application/octet-stream", None, "https://example.test/a.MP3", b"ID3-MP3", ".mp3", "audio/mpeg", "mp3"), + ("audio/x-wav", None, "https://example.test/a", b"RIFF-WAV", ".wav", "audio/wav", "wav"), + ("application/ogg", None, "https://example.test/a", b"OggS", ".ogg", "audio/ogg", None), + ("audio/ogg", None, "https://example.test/a", b"OggS", ".ogg", "audio/ogg", None), + ( + "application/octet-stream", + None, + "https://example.test/a", + b"UNKNOWN", + ".bin", + "application/octet-stream", + None, + ), + (None, "binary/octet-stream", "https://example.test/a", b"UNKNOWN", ".bin", "application/octet-stream", None), + ], +) +async def test_imported_audio_bytes_extension_and_downstream_format_async( + *, + sqlite_instance: SQLiteMemory, + content_type: str | None, + mime_type: str | None, + final_url: str, + content: bytes, + extension: str, + upload_type: str, + chat_format: str | None, +) -> None: + request_piece = MessagePieceRequest( + data_type="audio_path", original_value=final_url, mime_type=mime_type, import_url=True + ) + with _download(content_type=content_type, final_url=final_url, content=content): + await persist_message_pieces_async(pieces=[request_piece]) + + stored_path = Path(request_piece.original_value) + assert stored_path.suffix == extension + assert request_piece.data_type == "audio_path" + assert request_piece.converted_value == request_piece.original_value + async with aiofiles.open(stored_path, "rb") as stored_file: + assert await stored_file.read() == content + + piece = MessagePiece( + role="user", original_value=request_piece.original_value, original_value_data_type="audio_path" + ) + if chat_format is None: + with pytest.raises(ValueError, match="Unsupported audio format"): + await build_audio_content_entry_async(message_piece=piece) + else: + entry = await build_audio_content_entry_async(message_piece=piece) + assert entry["input_audio"]["format"] == chat_format + assert base64.b64decode(entry["input_audio"]["data"]) == content + + target = HTTPXAPITarget( + http_url="https://provider.test/upload", allowed_upload_directory=sqlite_instance.results_path + ) + client = MagicMock(spec=httpx.AsyncClient) + client.request = AsyncMock(return_value=httpx.Response(200, text="ok")) + with patch("pyrit.prompt_target.http_target.httpx_api_target.httpx.AsyncClient") as client_factory: + client_factory.return_value.__aenter__.return_value = client + await target._send_prompt_to_target_async(normalized_conversation=[Message(message_pieces=[piece])]) + assert client.request.call_args.kwargs["files"]["file"] == (stored_path.name, content, upload_type) + + @pytest.mark.parametrize( ("data_type", "content_type"), [ diff --git a/tests/unit/backend/test_media_url_import.py b/tests/unit/backend/test_media_url_import.py index 5c2b7c5379..7a92d690d7 100644 --- a/tests/unit/backend/test_media_url_import.py +++ b/tests/unit/backend/test_media_url_import.py @@ -135,27 +135,24 @@ def handler(request: httpx.Request) -> httpx.Response: assert not read_redirect_body -async def test_error_status_is_rejected_without_query_string( +async def test_error_status_preserves_exception_cause_async( transport: Callable[[Handler], list[httpx.Request]], caplog: pytest.LogCaptureFixture ) -> None: transport(lambda request: httpx.Response(404)) with ( - caplog.at_level(logging.WARNING, logger=media_url_import.__name__), caplog.at_level(logging.INFO, logger="httpx"), pytest.raises(ValueError, match="returned HTTP 404") as error, ): await download_media_url_async(url="https://user.example.test/cat.png?sv=1&sig=secret#frag") assert "secret" not in str(error.value) - assert error.value.__cause__ is None - assert error.value.__suppress_context__ + assert isinstance(error.value.__cause__, httpx.HTTPStatusError) + assert error.value.__cause__.response.status_code == 404 assert "https://user.example.test/cat.png" in str(error.value) - assert "HTTP 404 (HTTPStatusError)" in caplog.text - assert "HTTP Request: GET https://user.example.test/cat.png " in caplog.text - assert "secret" not in caplog.text + assert "HTTP Request: GET https://user.example.test/cat.png?sv=1&sig=secret#frag " in caplog.text -async def test_request_log_omits_query_string_only_while_downloading( +async def test_download_keeps_http_client_logs_unchanged_async( transport: Callable[[Handler], list[httpx.Request]], caplog: pytest.LogCaptureFixture ) -> None: trace_logger = logging.getLogger("httpcore.http11") @@ -172,8 +169,8 @@ def respond(request: httpx.Request) -> httpx.Response: async with httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(200))) as client: await client.get("https://other.test/page?keep=1") - assert "HTTP Request: GET https://example.test/cat.png " in caplog.text - assert "secret" not in caplog.text + assert "HTTP Request: GET https://example.test/cat.png?sig=secret " in caplog.text + assert "Location=https://example.test/next.png?sig=secret" in caplog.text assert "HTTP Request: GET https://other.test/page?keep=1 " in caplog.text assert "Location=https://other.test/next?keep=1" in caplog.text @@ -215,7 +212,6 @@ def respond(request: httpx.Request) -> httpx.Response: ) async def test_network_failures_name_the_reason_and_limit( transport: Callable[[Handler], list[httpx.Request]], - caplog: pytest.LogCaptureFixture, make_error: Callable[[httpx.Request], BaseException], reason: str, cause: str, @@ -225,17 +221,13 @@ def fail(request: httpx.Request) -> httpx.Response: transport(fail) - with ( - caplog.at_level(logging.WARNING, logger=media_url_import.__name__), - pytest.raises(ValueError) as error, - ): + with pytest.raises(ValueError) as error: await download_media_url_async(url="https://example.test/slow.png?sig=secret") assert str(error.value) == f"Media URL https://example.test/slow.png could not be downloaded: {reason}." - assert error.value.__cause__ is None - assert error.value.__suppress_context__ - assert f"https://example.test/slow.png could not be downloaded: {reason} ({cause}" in caplog.text - assert "secret" not in caplog.text + assert error.value.__cause__ is not None + assert type(error.value.__cause__).__name__ == cause + assert str(error.value.__cause__) == _SIGNED_DETAIL @pytest.mark.parametrize( diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 34b8ca8808..b81894fbb1 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -3212,7 +3212,7 @@ async def test_preparation_failure_uses_dispatch_stage_not_saved_messages_async( assert status.error @pytest.mark.parametrize("count", [1, 2]) - async def test_failed_url_import_logs_no_signed_query_async( + async def test_failed_url_import_preserves_exception_details_async( self, *, real_send_context: tuple[MessageSendService, AttackResult, MockPromptTarget, Base64Converter], @@ -3245,8 +3245,8 @@ def create_client() -> httpx.AsyncClient: assert status.failure_stage == MessageSendFailureStage.PREPARATION [failure] = [record for record in caplog.records if record.exc_info] assert "returned HTTP 404" in str(failure.exc_info[1]) - assert "secret" not in "".join(traceback.format_exception(*failure.exc_info)) - assert "secret" not in caplog.text + assert isinstance(failure.exc_info[1].__cause__, httpx.HTTPStatusError) + assert "https://example.test/cat.png?sig=secret" in "".join(traceback.format_exception(*failure.exc_info)) @pytest.mark.parametrize("count", [1, 3]) @pytest.mark.parametrize("failure", ["conversion", "normalization", "validation", "target", "metadata"])