diff --git a/.pyrit_conf_example b/.pyrit_conf_example index 0e1cdab013..2ac9b927d5 100644 --- a/.pyrit_conf_example +++ b/.pyrit_conf_example @@ -134,6 +134,22 @@ enable_live_reinitialization: false # Default: false allow_custom_initializers: false +# 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 9833855ea4..9994e3e17d 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -200,3 +200,49 @@ 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, 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 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 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 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. +- 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; 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: + +- 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. +- `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 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/routes/attacks.py b/pyrit/backend/routes/attacks.py index bb75f1f33b..0299f0c2ff 100644 --- a/pyrit/backend/routes/attacks.py +++ b/pyrit/backend/routes/attacks.py @@ -251,10 +251,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 14d4dfb2ec..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, @@ -787,6 +787,7 @@ async def _prepare_message_pieces_async( saved.original_value = request_piece.original_value converted_value = request_piece.converted_value saved.converted_value = converted_value if converted_value is not None else saved.original_value + 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 5ba393fc9d..0be70023b6 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -37,8 +37,13 @@ CreateConverterRequest, PreviewStep, ) -from pyrit.backend.services.media_persistence import persist_media_value_async -from pyrit.common.azure_storage import is_azure_blob_uri +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 from pyrit.models import MessagePiece, PromptDataType from pyrit.prompt_normalizer import ConverterConfiguration, PromptNormalizer @@ -216,15 +221,18 @@ 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, resolve references or persist base64/data URIs. + # 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, @@ -234,9 +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 + 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( @@ -248,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]: @@ -289,8 +300,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. + 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. @@ -308,7 +320,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 data URI or an http(s) URL, or a URL + cannot be downloaded. """ metadata = self._registry.get_registered_class_metadata(converter_type) path_params = ( @@ -330,13 +343,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): + 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 4985697c05..41349fa204 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -9,21 +9,39 @@ import base64 import binascii import mimetypes -from collections.abc import Callable, Sequence +from collections.abc import Callable, Coroutine, Mapping, 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.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 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"}) +# 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"}) +# 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): """Origin recognized for one path-typed media value.""" @@ -45,6 +63,7 @@ class MediaPersistenceResult: resolved: bool mime_type: str | None = None extension: str | None = None + source: str | None = None SerializerFactory = Callable[..., Any] @@ -66,6 +85,132 @@ 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 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. + + 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, @@ -85,6 +230,92 @@ def _resolve_extension( return extension or DEFAULT_MEDIA_EXTENSIONS.get(str(data_type), ".bin") +async def _write_owned_media_async( + *, + serializer: Any, + created_paths: list[str] | None, + write: Callable[[], Coroutine[Any, Any, None]], +) -> None: + """ + Write one media file, first recording it in ``created_paths`` when given. + + Ownership is recorded before writing so partial writes can also be removed, and a + cancelled request still waits for the write to finish. + """ + if created_paths is None: + await write() + return + created_paths.append(str(await serializer.get_data_filename_async())) + write_task = asyncio.create_task(write()) + try: + await asyncio.shield(write_task) + except asyncio.CancelledError: + await write_task + raise + + +async def _import_media_url_async( + *, + url: str, + data_type: PromptDataType, + mime_type: str | None, + serializer_factory: SerializerFactory, + created_paths: list[str] | None, +) -> MediaPersistenceResult: + """ + 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. 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. + + 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, 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 = media_extension(download, mime_type=mime_type, default=".bin") + serializer = serializer_factory( + category="prompt-memory-entries", + data_type=data_type, + extension=extension, + ) + await _write_owned_media_async( + serializer=serializer, + created_paths=created_paths, + write=lambda: 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, + 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, @@ -92,6 +323,7 @@ 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: @@ -100,12 +332,45 @@ 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 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. + 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 or must be imported. """ if value.startswith(("http://", "https://")): + if is_managed_blob_url(value): + return MediaPersistenceResult( + value=redact_url_credentials(value), + origin=MediaOrigin.REMOTE_URL, + persisted=False, + resolved=True, + mime_type=mime_type, + ) + if import_url: + return await _import_media_url_async( + 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( + 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, @@ -114,14 +379,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, ) @@ -133,17 +398,20 @@ async def persist_media_value_async( origin = MediaOrigin.DATA_URI else: try: - if await asyncio.to_thread(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, @@ -156,17 +424,11 @@ async def persist_media_value_async( data_type=data_type, extension=extension, ) - if created_paths is None: - await serializer.save_b64_image_async(data=payload) - else: - # Record ownership before writing so partial writes can also be removed. - created_paths.append(str(await serializer.get_data_filename_async())) - write_task = asyncio.create_task(serializer.save_b64_image_async(data=payload)) - try: - await asyncio.shield(write_task) - except asyncio.CancelledError: - await write_task - raise + await _write_owned_media_async( + serializer=serializer, + created_paths=created_paths, + write=lambda: serializer.save_b64_image_async(data=payload), + ) return MediaPersistenceResult( value=str(serializer.value), origin=origin, @@ -177,6 +439,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 resolved content type, 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], @@ -187,47 +479,65 @@ async def persist_message_pieces_async( 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. + 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. 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. + persisted_paths (list[str] | None): When given, receives the path of every file + written for these pieces, so a failed request can remove them. + 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 + 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 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, + import_url=piece.import_url, serializer_factory=serializer_factory, created_paths=persisted_paths, ) 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 - ): + source_metadata.update(media_source_metadata(result)) + if mirrors_original: 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) + 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 + source_metadata.update(media_source_metadata(result, prefix="converted_media_source")) piece.original_value = original_value piece.converted_value = converted_value + 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 new file mode 100644 index 0000000000..4f47342132 --- /dev/null +++ b/pyrit/backend/services/media_url_import.py @@ -0,0 +1,217 @@ +# 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 +_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}$") +_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 + + +@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, *, mime_type: str | None = None) -> str | None: + """ + Resolve the media type from the response, caller MIME type, then URL suffix. + + 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. + """ + 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, 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, 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 + 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.InvalidURL as exc: + raise ValueError(f"Media URL {shown} is not a valid http or https URL.") from exc + 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: + reason = _failure_reason(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. + + 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__})" diff --git a/pyrit/backend/services/runtime_lifecycle.py b/pyrit/backend/services/runtime_lifecycle.py index 75632819c4..94b8e44b58 100644 --- a/pyrit/backend/services/runtime_lifecycle.py +++ b/pyrit/backend/services/runtime_lifecycle.py @@ -14,12 +14,14 @@ 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, 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 @@ -101,6 +103,8 @@ 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) + 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 8d1176ed07..f09550edbf 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,28 @@ } +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 + + +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. @@ -184,15 +210,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 _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." + ) + async def list_target_types_async(self) -> TargetTypeResponse: """ List all available target types from the target class registry. - Returns every constructible target with its derived constructor - parameters and the auth modes it supports, all projected from the - registry's ``TargetMetadata``. Deciding which entries to surface to a - 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. @@ -201,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) @@ -220,11 +287,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 auth contract: 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. @@ -233,10 +302,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 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( @@ -244,6 +315,11 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn ) target_cls = self._registry.get_class(request.type) + 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": @@ -254,12 +330,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 cf04d20e52..11ae7eaabe 100644 --- a/pyrit/prompt_target/common/prompt_target.py +++ b/pyrit/prompt_target/common/prompt_target.py @@ -84,6 +84,15 @@ class PromptTarget(Identifiable): # Azure Blob Storage, Prompt Shield) override this to add ``"identity"``. supported_auth_modes: ClassVar[tuple[AuthMode, ...]] = ("api_key",) + # 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: """ 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..a5c90ae994 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,7 @@ class HTTPXAPITarget(HTTPTarget): """ _PATH_TYPES: frozenset[str] = frozenset({"image_path", "audio_path", "video_path", "binary_path"}) + 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 db37e5de81..904c522425 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. + loads_local_code: ClassVar[bool] = True + # Class-level cache for model and tokenizer _cached_model: Any = None _cached_tokenizer: Any = None diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index fb6d6850ed..c2559bc086 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -118,6 +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 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 @@ -162,6 +167,8 @@ 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 + 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) @@ -184,6 +191,12 @@ 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.") + 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_api_routes.py b/tests/unit/backend/test_api_routes.py index 2fe2d9286d..0e5e29179c 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -264,6 +264,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_service.py b/tests/unit/backend/test_attack_service.py index d915edad80..313d81345f 100644 --- a/tests/unit/backend/test_attack_service.py +++ b/tests/unit/backend/test_attack_service.py @@ -15,6 +15,7 @@ from pathlib import Path from typing import Any from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import quote import pytest from sqlalchemy import select @@ -43,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, @@ -919,6 +921,112 @@ async def test_create_attack_stores_prepended_conversation(self, attack_service, mock_memory.add_conversation_branches_to_attack_async.assert_called_once() assert len(mock_memory.add_conversation_branches_to_attack_async.call_args.kwargs["message_pieces"]) == 1 + 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_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) + ) + + mock_memory.add_conversation_branches_to_attack_async.assert_not_called() + 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.""" + stored_path = "/results/prompt-memory-entries/images/prepended.png" + serializer = MagicMock(value=stored_path) + serializer.get_data_filename_async = AsyncMock(return_value=stored_path) + 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, + 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) + ) + + assert factory.call_args.kwargs["category"] == "prompt-memory-entries" + 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: @@ -2110,17 +2218,22 @@ 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("reference", ["media_reference", "stored_file", "results_blob"]) async def test_converted_media_references_are_not_repersisted_async( - self, *, converted_value: str, expected_value: str + self, *, attack_service, mock_memory: MagicMock, tmp_path: Path, reference: str ) -> None: + stored_path = (tmp_path / "prompt-memory-entries" / "preview.png").resolve() + mock_memory.results_path = str(tmp_path) + if reference == "media_reference": + converted_value, expected_value = f"/api/media?path={stored_path}", str(stored_path) + elif reference == "stored_file": + stored_path.parent.mkdir(parents=True) + stored_path.write_bytes(b"PNG") + converted_value, expected_value = str(stored_path), str(stored_path) + else: + mock_memory.results_path = "https://account.blob.core.windows.net/results" + expected_value = f"{mock_memory.results_path}/prompt-memory-entries/preview.png" + converted_value = f"{expected_value}?sv=1" request = AddMessageRequest( pieces=[ MessagePieceRequest( @@ -2132,10 +2245,7 @@ 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.attack_service.data_serializer_factory") as factory, - ): + with patch("pyrit.backend.services.attack_service.data_serializer_factory") as factory: await AttackService._persist_base64_pieces_async(pieces=request.pieces) assert request.pieces[0].original_value == "source" @@ -2381,14 +2491,16 @@ async def test_path_data_type_supplies_extension_when_mime_type_missing(self, at ) assert request.pieces[0].original_value == "/saved/image.png" - async def test_http_url_is_kept_as_is(self, attack_service) -> None: - """HTTPS blob URLs should not be re-persisted.""" + async def test_http_url_is_kept_as_is(self, attack_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" 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=f"{blob_url}?sv=2024", mime_type="image/png", ), ], @@ -2398,17 +2510,21 @@ async def test_http_url_is_kept_as_is(self, attack_service) -> None: await AttackService._persist_base64_pieces_async(pieces=request.pieces) - 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, attack_service) -> None: - """Local media URLs are converted back to their decoded file paths.""" + async def test_media_reference_is_resolved_without_persistence( + self, attack_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={quote(str(stored_path), safe='')}", ), ], send=False, @@ -2418,13 +2534,17 @@ async def test_media_reference_is_resolved_without_persistence(self, attack_serv with patch("pyrit.backend.services.attack_service.data_serializer_factory") as factory: await AttackService._persist_base64_pieces_async(pieces=request.pieces) - 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, attack_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, attack_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").resolve() + media_path.parent.mkdir(parents=True) media_path.write_bytes(b"image") request = AddMessageRequest( role="user", diff --git a/tests/unit/backend/test_conversation_editor.py b/tests/unit/backend/test_conversation_editor.py index 68732c6a24..817ffb733d 100644 --- a/tests/unit/backend/test_conversation_editor.py +++ b/tests/unit/backend/test_conversation_editor.py @@ -153,11 +153,11 @@ async def test_targetless_media_is_saved_but_incompatible_binding_writes_nothing *, response_target: OpenAIResponseTarget, sqlite_instance: SQLiteMemory, - tmp_path: Path, data_type: PromptDataType, related: bool, ) -> None: - media = tmp_path / "history.bin" + media = Path(sqlite_instance.results_path) / "prompt-memory-entries" / "history.bin" + await asyncio.to_thread(media.parent.mkdir, parents=True, exist_ok=True) await asyncio.to_thread(media.write_bytes, b"history bytes") service = AttackService() source = await service.save_conversation_async(request=draft()) @@ -211,9 +211,9 @@ async def test_supported_converted_history_can_save_and_bind_async( *, response_target: OpenAIResponseTarget, sqlite_instance: SQLiteMemory, - tmp_path: Path, ) -> None: - media = tmp_path / "original.wav" + media = Path(sqlite_instance.results_path) / "prompt-memory-entries" / "original.wav" + await asyncio.to_thread(media.parent.mkdir, parents=True, exist_ok=True) await asyncio.to_thread(media.write_bytes, b"original audio bytes") service = AttackService() request = draft() diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index ed39ed57d7..be0c4164d4 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -26,6 +26,7 @@ ConverterService, get_converter_service, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.converter import ( Base64Converter, BinaryConverter, @@ -44,6 +45,14 @@ from pyrit.registry.components import ConverterRegistry from unit.mocks import MockPromptTarget +_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", @@ -606,10 +615,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 @@ -618,11 +628,69 @@ 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("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( + "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", + "https://example.org/input.mp4", + ], + ) + async def test_path_or_str_url_outside_results_is_downloaded_to_owned_upload( + self, upload_service: ConverterService, url: str + ) -> None: + 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": "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", "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} @@ -1015,31 +1083,114 @@ 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() + + 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="/api/media?path=/etc/hostname", original_value_data_type="image_path", converter_ids=[] + ) + + 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_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() + 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?sig=secret", + original_value_data_type="image_path", + converter_ids=[], + import_url=True, + ) + + 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?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, 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.""" service = ConverterService() @@ -1078,7 +1229,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( @@ -1105,6 +1256,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, @@ -1294,11 +1446,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") @@ -1443,9 +1598,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), @@ -1453,11 +1609,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..19f680b4a8 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -3,56 +3,393 @@ """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 + +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, 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 -from pyrit.backend.services.media_persistence import MediaOrigin, persist_media_value_async +_BLOB_ROOT = "https://account.blob.core.windows.net/results" 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 +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" + 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_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 == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" + 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://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_existing_references_are_not_persisted( - value: str, origin: MediaOrigin, resolved_value: str, resolved: bool -) -> None: +async def test_url_is_kept_as_a_reference_without_import(value: str) -> None: factory = MagicMock() - 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", serializer_factory=factory) - assert result.origin is origin - assert result.value == resolved_value - assert result.resolved is resolved + 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(_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") + 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.source is not None + assert "?" not in result.source + + +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", import_url=True, serializer_factory=MagicMock() + ) + + download.assert_not_awaited() + assert result.value == f"{_BLOB_ROOT}/prompt-memory-entries/images/stored.png" + assert result.persisted is False + + +@pytest.mark.parametrize( + ("data_type", "content_type", "final_url", "expected_extension"), + [ + ("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", ".ogg"), + ("binary_path", "application/pdf", "https://example.test/a", ".pdf"), + ("binary_path", "image/png", "https://example.test/a", ".png"), + ], +) +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): + 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=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"), + [ + ("image_path", "text/html"), + ("image_path", "audio/wav"), + ("audio_path", "image/png"), + ("video_path", "text/plain"), + ], +) +async def test_imported_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, + import_url=True, + serializer_factory=factory, + ) + + assert "secret" not in str(error.value) 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", + [ + 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) @@ -142,3 +479,91 @@ 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_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, 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="image_path", + import_url=True, + ) + 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") + assert piece.prompt_metadata == { + "converted_media_source_url": "https://example.test/cat", + "converted_media_source_content_type": "image/png", + } + + +@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 new file mode 100644 index 0000000000..7a92d690d7 --- /dev/null +++ b/tests/unit/backend/test_media_url_import.py @@ -0,0 +1,282 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for the bounded media URL download.""" + +import logging +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_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.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 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 Request: GET https://user.example.test/cat.png?sv=1&sig=secret#frag " in caplog.text + + +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") + + 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?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 + + +_SIGNED_DETAIL = "failed for https://example.test/slow.png?sig=secret" + + +@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]], + make_error: Callable[[httpx.Request], BaseException], + reason: str, + cause: str, +) -> None: + def fail(request: httpx.Request) -> httpx.Response: + raise make_error(request) + + transport(fail) + + 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 not None + assert type(error.value.__cause__).__name__ == cause + assert str(error.value.__cause__) == _SIGNED_DETAIL + + +@pytest.mark.parametrize( + ("url", "message"), + [ + ("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"), + ], +) +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 1ac6b74247..b81894fbb1 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -4,6 +4,8 @@ """Tests for the shared synchronous manual-message owner.""" import asyncio +import logging +import traceback import uuid from collections.abc import AsyncGenerator, Generator, Iterator, Sequence from contextlib import asynccontextmanager, contextmanager @@ -12,6 +14,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 @@ -36,6 +39,7 @@ ManualSendScheduler, get_manual_send_scheduler, ) +from pyrit.backend.services.media_url_import import MediaDownload from pyrit.backend.services.message_send_service import ( MessageSendNotFoundError, MessageSendService, @@ -77,6 +81,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()) @@ -473,6 +486,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 = [] @@ -531,7 +545,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", @@ -807,17 +821,18 @@ 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" + 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) request = AddMessageRequest( pieces=[ MessagePieceRequest( @@ -829,16 +844,60 @@ 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() + async def test_converted_media_outside_results_is_rejected_async( + self, *, mock_memory: MagicMock, tmp_path: Path + ) -> None: + mock_memory.results_path = str(tmp_path) + request = AddMessageRequest( + pieces=[ + MessagePieceRequest( + original_value="source", + converted_value="/api/media?path=/etc/hostname", + converted_value_data_type="image_path", + ) + ], + send=False, + target_conversation_id="test-id", + ) + + with pytest.raises(ValueError, match="results directory"): + await MessageSendService._persist_base64_pieces_async(request) + + 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="image_path", original_value="https://example.com/cat?sig=secret", import_url=True + ) + ], + 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=[ @@ -894,11 +953,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=[ @@ -1083,14 +1145,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_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" 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=f"{blob_url}?sv=2024&sig=secret", mime_type="image/png", ), ], @@ -1100,17 +1164,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, @@ -1120,13 +1188,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", @@ -1138,10 +1210,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( @@ -2575,22 +2664,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", ), ], ) @@ -2599,6 +2688,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, @@ -2606,6 +2696,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( @@ -3111,6 +3211,43 @@ 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_preserves_exception_details_async( + self, + *, + real_send_context: tuple[MessageSendService, AttackResult, MockPromptTarget, Base64Converter], + caplog: pytest.LogCaptureFixture, + count: int, + ) -> None: + service, ar, _, _ = real_send_context + caplog.set_level(logging.INFO) + 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 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"]) async def test_real_sqlite_failure_stages_preserve_evidence_async( diff --git a/tests/unit/backend/test_runtime_lifecycle.py b/tests/unit/backend/test_runtime_lifecycle.py index 3519a1819d..a6cf728b64 100644 --- a/tests/unit/backend/test_runtime_lifecycle.py +++ b/tests/unit/backend/test_runtime_lifecycle.py @@ -330,6 +330,32 @@ 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) + + +@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 e6803cc181..f0a18be66d 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 @@ -263,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"}), ], @@ -299,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() @@ -405,6 +424,59 @@ 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", "upload_directory", "error"), + [ + ("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_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(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 04b0b3f39d..a91eb700a8 100644 --- a/tests/unit/setup/test_configuration_loader.py +++ b/tests/unit/setup/test_configuration_loader.py @@ -76,6 +76,30 @@ 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_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] + + @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"]: