Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 9 additions & 50 deletions pyrit/backend/routes/media.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,22 +16,19 @@
opaque bytes.
"""

import asyncio
import logging
import mimetypes
from pathlib import Path

from fastapi import APIRouter, HTTPException, Query
from fastapi.responses import FileResponse

from pyrit.memory import CentralMemory
from pyrit.backend.services.media_persistence import MediaAccessDeniedError, validate_local_media_path_async

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 = {
Expand Down Expand Up @@ -60,41 +57,6 @@
}


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``).

Args:
path: The user-provided file path to validate.
allowed_root: The canonical (``resolve``-d) allowed root directory.

Returns:
The canonical, validated file 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
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


@router.get("/media")
async def serve_media_async(
path: str = Query(..., description="Absolute path to the local media file to serve."),
Expand All @@ -120,19 +82,16 @@ async def serve_media_async(
Raises:
HTTPException 403: If the path is outside the allowed directory.
HTTPException 404: If the file does not exist.
HTTPException 500: If memory is not initialized.
HTTPException 500: If memory or its results path is not configured.
"""
try:
memory = CentralMemory.get_memory_instance()
if not memory.results_path:
raise HTTPException(status_code=500, detail="Memory results_path is not configured.")
allowed_root = Path(memory.results_path).resolve(strict=False)
except Exception as exc:
raise HTTPException(status_code=500, detail="Memory not initialized; cannot determine results path.") from exc

validated_path = _validate_media_path(path=path, allowed_root=allowed_root)
validated_path = await validate_local_media_path_async(path=path)
except MediaAccessDeniedError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
except RuntimeError as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc

if not validated_path.is_file():
if not await asyncio.to_thread(validated_path.is_file):
raise HTTPException(status_code=404, detail="File not found.")

extension = validated_path.suffix.lower()
Expand Down
98 changes: 81 additions & 17 deletions pyrit/backend/services/media_persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,21 @@
from urllib.parse import parse_qs, urlparse

from pyrit.backend.models import DEFAULT_MEDIA_EXTENSIONS
from pyrit.memory import data_serializer_factory
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


_ALLOWED_MEDIA_SUBDIRECTORIES = frozenset({"prompt-memory-entries", "seed-prompt-entries"})


class MediaAccessDeniedError(ValueError):
"""A local media path is outside the API's allowed storage directories."""


class MediaOrigin(str, Enum):
"""Origin recognized for one path-typed media value."""

Expand All @@ -50,6 +57,55 @@ class MediaPersistenceResult:
SerializerFactory = Callable[..., Any]


def validate_media_path(*, path: str, allowed_root: Path) -> Path:
"""
Resolve symlinks and parent components before checking results-directory containment.

Args:
path (str): The local media path.
allowed_root (Path): The configured results directory.

Returns:
The canonical path in an allowed media subdirectory.

Raises:
MediaAccessDeniedError: If the path is outside the allowed media directories.
"""
real_path = Path(path).resolve(strict=False)
try:
relative_parts = real_path.relative_to(allowed_root.resolve(strict=False)).parts
except ValueError as exc:
raise MediaAccessDeniedError("Access denied: path is outside the allowed results directory.") from exc

if not relative_parts or relative_parts[0] not in _ALLOWED_MEDIA_SUBDIRECTORIES:
raise MediaAccessDeniedError("Access denied: path is not in a media subdirectory.")

return real_path


async def validate_local_media_path_async(*, path: str) -> Path:
"""
Validate a local media path against the configured results directory.

Args:
path (str): The local media path.

Returns:
The canonical path in an allowed media subdirectory.

Raises:
RuntimeError: If memory or its results path is not configured.
MediaAccessDeniedError: If the path is outside the allowed media directories.
"""
try:
memory = CentralMemory.get_memory_instance()
except ValueError as exc:
raise RuntimeError("Memory not initialized; cannot determine results path.") from exc
if not memory.results_path:
raise RuntimeError("Memory results_path is not configured.")
return await asyncio.to_thread(validate_media_path, path=path, allowed_root=Path(memory.results_path))


def _is_raw_base64(value: str) -> bool:
"""Return whether *value* is syntactically valid raw base64."""
try:
Expand Down Expand Up @@ -102,6 +158,9 @@ async def persist_media_value_async(
attack ingestion and converter preview while keeping origin detection,
extension resolution, and persistence in one component.

Local paths and ``/api/media`` references must be in a media subdirectory
of the configured results directory, as required by ``GET /api/media``.

Returns:
A typed result containing the resolved value and persistence metadata.
"""
Expand All @@ -116,14 +175,17 @@ async def persist_media_value_async(

if value.startswith("/api/media"):
parsed = urlparse(value)
file_path = parse_qs(parsed.query).get("path", [None])[0]
return MediaPersistenceResult(
value=file_path or value,
origin=MediaOrigin.MEDIA_REFERENCE,
persisted=False,
resolved=file_path is not None,
mime_type=mime_type,
)
if parsed.path == "/api/media":
file_path = parse_qs(parsed.query).get("path", [None])[0]
if file_path is None:
raise ValueError("Media reference must include a path.")
return MediaPersistenceResult(
value=str(await validate_local_media_path_async(path=file_path)),
origin=MediaOrigin.MEDIA_REFERENCE,
persisted=False,
resolved=True,
mime_type=mime_type,
)

data_uri_mime_type: str | None = None
payload = value
Expand All @@ -133,17 +195,19 @@ 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_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_file = False
if is_file:
return MediaPersistenceResult(
value=str(await validate_local_media_path_async(path=value)),
origin=MediaOrigin.LOCAL_PATH,
persisted=False,
resolved=True,
mime_type=mime_type,
)

extension = _resolve_extension(
data_type=data_type,
Expand Down
14 changes: 13 additions & 1 deletion tests/unit/backend/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@
"""Backend compatibility fixtures independent of a packaged workspace stamp."""

from collections.abc import Iterator
from unittest.mock import patch
from pathlib import Path
from unittest.mock import MagicMock, patch

import pytest

Expand All @@ -13,6 +14,17 @@
from pyrit.backend.services.attack_service import get_attack_service
from pyrit.backend.services.manual_send_scheduler import get_manual_send_scheduler
from pyrit.backend.services.message_send_service import get_message_send_service
from pyrit.memory import SQLiteMemory


@pytest.fixture
def managed_media_path(*, sqlite_instance: SQLiteMemory, patch_central_database: MagicMock) -> Path:
"""Create a file in the isolated memory's media directory."""
assert sqlite_instance.results_path is not None
path = Path(sqlite_instance.results_path) / "prompt-memory-entries" / "image.png"
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(b"image")
return path


@pytest.fixture(autouse=True)
Expand Down
61 changes: 41 additions & 20 deletions tests/unit/backend/test_attack_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from pathlib import Path
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
from urllib.parse import urlencode

import pytest
from sqlalchemy import select
Expand Down Expand Up @@ -813,6 +814,33 @@ async def test_get_conversation_messages_raises_for_unrelated_conversation(
class TestCreateAttack:
"""Tests for create_attack method."""

@pytest.mark.parametrize("reference", [False, True])
@pytest.mark.parametrize("field", ["original_value", "converted_value"])
async def test_rejects_unmanaged_prepended_media_async(
self, *, managed_media_path: Path, tmp_path: Path, reference: bool, field: str
) -> None:
path = tmp_path / "outside.png"
path.write_bytes(b"image")
value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path)
piece = MessagePieceRequest(
data_type="image_path",
original_value=value if field == "original_value" else str(managed_media_path),
converted_value=value if field == "converted_value" else str(managed_media_path),
)
service = AttackService()
with (
patch.object(service, "_get_save_target_async", new=AsyncMock(return_value=None)),
patch.object(service._memory, "add_conversation_branches_to_attack_async", new_callable=AsyncMock) as store,
pytest.raises(ValueError, match="outside the allowed results directory"),
):
await service.create_attack_async(
request=CreateAttackRequest(
prepended_conversation=[PrependedMessageRequest(role="user", pieces=[piece])]
)
)

store.assert_not_awaited()

@pytest.mark.parametrize("copy_history", [False, True])
async def test_manual_creation_registers_ownership_async(
self, *, sqlite_instance: SQLiteMemory, copy_history: bool
Expand Down Expand Up @@ -2110,17 +2138,14 @@ 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("origin", ["reference", "url", "local"])
async def test_converted_media_references_are_not_repersisted_async(
self, *, converted_value: str, expected_value: str
self, *, origin: str, managed_media_path: Path
) -> None:
expected_value = "https://example.com/preview.png?token=example" if origin == "url" else str(managed_media_path)
converted_value = (
f"/api/media?{urlencode({'path': expected_value})}" if origin == "reference" else expected_value
)
request = AddMessageRequest(
pieces=[
MessagePieceRequest(
Expand All @@ -2132,10 +2157,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"
Expand Down Expand Up @@ -2401,14 +2423,14 @@ async def test_http_url_is_kept_as_is(self, attack_service) -> None:
assert request.pieces[0].original_value == ("https://myblob.blob.core.windows.net/images/photo.png?sv=2024")
assert request.pieces[0].converted_value == request.pieces[0].original_value

async def test_media_reference_is_resolved_without_persistence(self, attack_service) -> None:
async def test_media_reference_is_resolved_without_persistence(self, *, managed_media_path: Path) -> None:
"""Local media URLs are converted back to their decoded file paths."""
request = AddMessageRequest(
role="user",
pieces=[
MessagePieceRequest(
data_type="image_path",
original_value="/api/media?path=%2Ftmp%2Fimage.png",
original_value=f"/api/media?{urlencode({'path': str(managed_media_path)})}",
),
],
send=False,
Expand All @@ -2418,14 +2440,13 @@ 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(managed_media_path)
assert request.pieces[0].converted_value == str(managed_media_path)
factory.assert_not_called()

async def test_existing_file_is_kept_without_persistence(self, attack_service, tmp_path: Path) -> None:
async def test_existing_file_is_kept_without_persistence(self, *, managed_media_path: Path) -> None:
"""An existing path remains the canonical original and converted value."""
media_path = tmp_path / "image.png"
media_path.write_bytes(b"image")
media_path = managed_media_path
request = AddMessageRequest(
role="user",
pieces=[MessagePieceRequest(data_type="image_path", original_value=str(media_path))],
Expand Down
Loading
Loading