From dcfe920b7eea2a95ad34ca6a39be42f23f1832a3 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 9 Oct 2026 10:13:52 -0700 Subject: [PATCH 1/2] Fix API local media path containment Reuse the media route's canonical containment check for preview and message inputs. Validate original and converted media, reject incomplete media references, and keep URLs, uploads, and converter outputs unchanged. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 59fd033b-89ce-4424-b418-510b8144b8c8 --- pyrit/backend/routes/media.py | 49 +----- pyrit/backend/services/media_persistence.py | 72 +++++++-- tests/unit/backend/conftest.py | 14 +- tests/unit/backend/test_attack_service.py | 61 ++++--- tests/unit/backend/test_converter_service.py | 43 ++++- tests/unit/backend/test_media_persistence.py | 151 +++++++++++++++++- .../unit/backend/test_message_send_service.py | 86 +++++++--- 7 files changed, 363 insertions(+), 113 deletions(-) diff --git a/pyrit/backend/routes/media.py b/pyrit/backend/routes/media.py index da93d1e3e1..0d6b46ffbd 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -16,6 +16,7 @@ opaque bytes. """ +import asyncio import logging import mimetypes from pathlib import Path @@ -23,15 +24,13 @@ from fastapi import APIRouter, HTTPException, Query from fastapi.responses import FileResponse +from pyrit.backend.services.media_persistence import validate_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 = { @@ -60,41 +59,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."), @@ -126,13 +90,16 @@ async def serve_media_async( 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) + allowed_root = await asyncio.to_thread(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) + try: + validated_path = await asyncio.to_thread(validate_media_path, path=path, allowed_root=allowed_root) + except ValueError as exc: + raise HTTPException(status_code=403, 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() diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index 4985697c05..66deb83a05 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -17,7 +17,7 @@ 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: @@ -50,6 +50,36 @@ 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. + + Returns: + The canonical path in an allowed media subdirectory. + + Raises: + ValueError: 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 ValueError("Access denied: path is outside the allowed results directory.") from exc + + if not relative_parts or relative_parts[0] not in {"prompt-memory-entries", "seed-prompt-entries"}: + raise ValueError("Access denied: path is not in a media subdirectory.") + + return real_path + + +async def _validate_local_media_path_async(path: str) -> str: + memory = CentralMemory.get_memory_instance() + if not memory.results_path: + raise ValueError("Memory results_path is not configured.") + validated_path = await asyncio.to_thread(validate_media_path, path=path, allowed_root=Path(memory.results_path)) + return str(validated_path) + + def _is_raw_base64(value: str) -> bool: """Return whether *value* is syntactically valid raw base64.""" try: @@ -102,6 +132,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. """ @@ -116,14 +149,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=await _validate_local_media_path_async(file_path), + origin=MediaOrigin.MEDIA_REFERENCE, + persisted=False, + resolved=True, + mime_type=mime_type, + ) data_uri_mime_type: str | None = None payload = value @@ -133,17 +169,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=await _validate_local_media_path_async(value), + origin=MediaOrigin.LOCAL_PATH, + persisted=False, + resolved=True, + mime_type=mime_type, + ) extension = _resolve_extension( data_type=data_type, diff --git a/tests/unit/backend/conftest.py b/tests/unit/backend/conftest.py index b224374953..9572992b52 100644 --- a/tests/unit/backend/conftest.py +++ b/tests/unit/backend/conftest.py @@ -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 @@ -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) diff --git a/tests/unit/backend/test_attack_service.py b/tests/unit/backend/test_attack_service.py index d915edad80..9fa7aa4826 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 urlencode import pytest from sqlalchemy import select @@ -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 @@ -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( @@ -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" @@ -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, @@ -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))], diff --git a/tests/unit/backend/test_converter_service.py b/tests/unit/backend/test_converter_service.py index ed39ed57d7..65560cb681 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -11,6 +11,7 @@ from collections.abc import AsyncGenerator from pathlib import Path from unittest.mock import AsyncMock, MagicMock, call, patch +from urllib.parse import quote, urlencode import pytest from fastapi import HTTPException @@ -1019,14 +1020,16 @@ async def test_preview_conversion_with_converter_ids(self) -> None: ("value", "resolved_value"), [ ("https://example.test/image.png", "https://example.test/image.png"), - ("/api/media?path=%2Ftmp%2Fimage.png", "/tmp/image.png"), + ("/api/media?path={path}", "{path}"), ], ) async def test_preview_conversion_resolves_reference_without_persistence( - self, value: str, resolved_value: str + self, *, value: str, resolved_value: str, managed_media_path: Path ) -> None: """Remote and local media references bypass serializer persistence.""" service = ConverterService() + value = value.format(path=quote(str(managed_media_path), safe="")) + resolved_value = resolved_value.format(path=str(managed_media_path)) request = ConverterPreviewRequest( original_value=value, original_value_data_type="image_path", @@ -1289,8 +1292,10 @@ async def test_preview_conversion_rejects_invalid_selection_before_conversion_as convert.assert_not_awaited() async def test_preview_conversion_unmarked_media_retains_result_type_async( - self, upload_service: ConverterService + self, *, upload_service: ConverterService, tmp_path: Path ) -> None: + output_path = tmp_path / "converted.wav" + output_path.write_bytes(b"RIFF") instance = Base64Converter() upload_service._registry.instances.register(instance, name="media") request = ConverterPreviewRequest( @@ -1299,10 +1304,10 @@ async def test_preview_conversion_unmarked_media_retains_result_type_async( converter_ids=["media"], ) with patch.object(instance, "convert_async", new_callable=AsyncMock) as convert: - convert.return_value = converter.ConverterResult(output_text="converted.wav", output_type="audio_path") + convert.return_value = converter.ConverterResult(output_text=str(output_path), 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") - assert result.converted_value == "converted.wav" + assert result.converted_value == str(output_path) assert result.converted_value_data_type == "audio_path" assert result.steps[0].input_data_type == "image_path" assert result.steps[0].output_data_type == "audio_path" @@ -1442,11 +1447,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: + async def test_preview_conversion_preserves_existing_file(self, *, managed_media_path: Path) -> None: """Existing local media paths pass through without being persisted again.""" service = ConverterService() - media_path = tmp_path / "input.wav" - media_path.write_bytes(b"RIFF") + media_path = managed_media_path request = ConverterPreviewRequest( original_value=str(media_path), original_value_data_type="audio_path", @@ -1459,6 +1463,29 @@ async def test_preview_conversion_preserves_existing_file(self, tmp_path: Path) mock_factory.assert_not_called() assert result.converted_value == str(media_path) + @pytest.mark.parametrize("reference", [False, True]) + async def test_preview_rejects_unmanaged_file_before_conversion_async( + self, *, upload_service: ConverterService, tmp_path: Path, reference: bool + ) -> None: + path = tmp_path / "outside.png" + path.write_bytes(b"image") + value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + instance = Base64Converter() + upload_service._registry.instances.register(instance, name="base64") + request = ConverterPreviewRequest( + original_value=value, original_value_data_type="image_path", converter_ids=["base64"] + ) + with ( + patch.object(instance, "convert_async", new_callable=AsyncMock) as convert, + patch.object(converter_routes, "get_converter_service", return_value=upload_service), + pytest.raises(HTTPException) as exc_info, + ): + await converter_routes.preview_conversion(request) + + assert exc_info.value.status_code == 400 + assert "outside the allowed results directory" in exc_info.value.detail + convert.assert_not_awaited() + class TestGetConverterObjectsForIds: """Tests for ConverterService.get_converter_objects_for_ids method.""" diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index 2fdab2f1af..e4c597d1a3 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -5,10 +5,17 @@ from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import urlencode import pytest -from pyrit.backend.services.media_persistence import MediaOrigin, persist_media_value_async +from pyrit.backend.models.attacks import MessagePieceRequest +from pyrit.backend.services.media_persistence import ( + MediaOrigin, + persist_media_value_async, + persist_message_pieces_async, +) +from pyrit.memory import SQLiteMemory def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: @@ -22,8 +29,7 @@ def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: ("value", "origin", "resolved_value", "resolved"), [ ("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), + ("http://example.test/media.png", MediaOrigin.REMOTE_URL, "http://example.test/media.png", True), ], ) async def test_existing_references_are_not_persisted( @@ -40,9 +46,18 @@ async def test_existing_references_are_not_persisted( 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", ["/api/media", "/api/media?path=", "/api/media?other=image.png"]) +async def test_media_reference_requires_path_async(value: str) -> None: + factory = MagicMock() + + with pytest.raises(ValueError, match="Media reference must include a path"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + + factory.assert_not_called() + + +async def test_existing_local_path_is_not_persisted_async(managed_media_path: Path) -> None: + media_path = managed_media_path factory = MagicMock() result = await persist_media_value_async(value=str(media_path), data_type="audio_path", serializer_factory=factory) @@ -53,6 +68,130 @@ async def test_existing_local_path_is_not_persisted(tmp_path: Path) -> None: factory.assert_not_called() +@pytest.mark.usefixtures("patch_central_database") +class TestLocalMediaPaths: + async def test_media_prefix_does_not_bypass_local_path_check_async(self) -> None: + with ( + patch.object(Path, "is_file", return_value=True), + pytest.raises(ValueError, match="outside the allowed results directory"), + ): + await persist_media_value_async(value="/api/media-other.png", data_type="image_path") + + @pytest.mark.parametrize("reference", [False, True]) + @pytest.mark.parametrize("folder", ["prompt-memory-entries", "seed-prompt-entries"]) + async def test_accepts_allowed_media_folders_async( + self, *, managed_media_path: Path, reference: bool, folder: str + ) -> None: + path = managed_media_path.parent.parent / folder / "nested" / "image.png" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(b"image") + value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + factory = MagicMock() + + result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + + assert result.value == str(path.resolve()) + assert result.origin is (MediaOrigin.MEDIA_REFERENCE if reference else MediaOrigin.LOCAL_PATH) + assert result.resolved is True + assert result.persisted is False + factory.assert_not_called() + + @pytest.mark.parametrize("reference", [False, True]) + @pytest.mark.parametrize("location", ["outside", "results-root", "other-directory", "sibling", "traversal"]) + async def test_rejects_unmanaged_files_async( + self, *, managed_media_path: Path, tmp_path: Path, reference: bool, location: str + ) -> None: + root = managed_media_path.parent.parent + paths = { + "outside": tmp_path / "outside.png", + "results-root": root / "image.png", + "other-directory": root / "other" / "image.png", + "sibling": root / "prompt-memory-entries-other" / "image.png", + "traversal": managed_media_path.parent / ".." / "image.png", + } + path = paths[location] + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(b"image") + value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + factory = MagicMock() + + with pytest.raises(ValueError, match="Access denied"): + await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + + factory.assert_not_called() + + @pytest.mark.parametrize("reference", [False, True]) + async def test_returns_canonical_allowed_path_async(self, *, managed_media_path: Path, reference: bool) -> None: + nested = managed_media_path.parent / "nested" + nested.mkdir() + path = nested / ".." / managed_media_path.name + value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + + result = await persist_media_value_async(value=value, data_type="image_path") + + assert result.value == str(managed_media_path.resolve()) + + @pytest.mark.parametrize("reference", [False, True]) + async def test_checks_resolved_path_not_input_path_async( + self, *, managed_media_path: Path, tmp_path: Path, reference: bool + ) -> None: + outside = tmp_path / "outside.png" + value = f"/api/media?{urlencode({'path': str(managed_media_path)})}" if reference else str(managed_media_path) + with ( + patch.object(Path, "resolve", side_effect=[outside, managed_media_path.parent.parent]), + pytest.raises(ValueError, match="outside the allowed results directory"), + ): + await persist_media_value_async(value=value, data_type="image_path") + + @pytest.mark.parametrize("reference", [False, True]) + async def test_rejects_symlink_escape_async( + self, *, managed_media_path: Path, tmp_path: Path, reference: bool + ) -> None: + outside = tmp_path / "outside.png" + outside.write_bytes(b"image") + path = managed_media_path.parent / "link.png" + try: + path.symlink_to(outside) + except (OSError, NotImplementedError) as exc: + pytest.skip(f"Cannot create symlink in this environment: {exc}") + value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + + with pytest.raises(ValueError, match="outside the allowed results directory"): + await persist_media_value_async(value=value, data_type="image_path") + + @pytest.mark.parametrize("reference", [False, True]) + @pytest.mark.parametrize("field", ["original_value", "converted_value"]) + async def test_checks_both_message_values_async( + self, *, managed_media_path: Path, tmp_path: Path, reference: bool, field: str + ) -> None: + outside = tmp_path / "outside.png" + outside.write_bytes(b"image") + value = f"/api/media?{urlencode({'path': str(outside)})}" if reference else str(outside) + 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), + converted_value_data_type="audio_path", + ) + before = piece.model_dump() + factory = MagicMock() + + with pytest.raises(ValueError, match="outside the allowed results directory"): + await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) + + assert piece.model_dump() == before + factory.assert_not_called() + + async def test_requires_configured_results_path_async( + self, *, managed_media_path: Path, sqlite_instance: SQLiteMemory + ) -> None: + with ( + patch.object(sqlite_instance, "results_path", None), + pytest.raises(ValueError, match="results_path is not configured"), + ): + await persist_media_value_async(value=str(managed_media_path), data_type="image_path") + + async def test_data_uri_uses_explicit_mime_before_uri_mime() -> None: serializer = _serializer(value="/saved/media.jpg") factory = MagicMock(return_value=serializer) diff --git a/tests/unit/backend/test_message_send_service.py b/tests/unit/backend/test_message_send_service.py index 1ac6b74247..94ee05e37f 100644 --- a/tests/unit/backend/test_message_send_service.py +++ b/tests/unit/backend/test_message_send_service.py @@ -11,6 +11,7 @@ from pathlib import Path from typing import Any from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import urlencode import pytest from pydantic import ValidationError @@ -807,17 +808,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( @@ -829,10 +827,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.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" @@ -1103,14 +1098,14 @@ async def test_http_url_is_kept_as_is(self, message_send_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, message_send_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, @@ -1120,14 +1115,13 @@ 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(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, message_send_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))], @@ -2605,7 +2599,19 @@ async def test_send_preserves_exact_preview_and_converts_other_piece_async( converted_value: str, expected_original: str, expected_final: str, + managed_media_path: Path, ) -> None: + mock_memory.results_path = str(managed_media_path.parent.parent) + if original_type.endswith("_path"): + original_value = f"/api/media?{urlencode({'path': str(managed_media_path)})}" + expected_original = str(managed_media_path) + if converted_type.endswith("_path"): + converted_path = managed_media_path.with_name( + "preview.wav" if converted_type == "audio_path" else "preview.png" + ) + converted_path.write_bytes(b"preview") + converted_value = f"/api/media?{urlencode({'path': str(converted_path)})}" + expected_final = str(converted_path) 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( @@ -2715,6 +2721,46 @@ def test_preconverted_provenance_preserves_explicit_execution_order(self) -> Non ] +@pytest.mark.usefixtures("patch_central_database") +class TestLocalMediaInputs: + @pytest.mark.parametrize("reference", [False, True]) + @pytest.mark.parametrize("field", ["original_value", "converted_value"]) + async def test_rejects_unmanaged_media_before_send_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) + request = AddMessageRequest( + target_conversation_id="test-id", + pieces=[ + 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 = MessageSendService(scheduler=ManualSendScheduler()) + target = MockPromptTarget() + with ( + patch.object(target, "send_prompt_async", new_callable=AsyncMock) as send, + pytest.raises(ValueError, match="outside the allowed results directory"), + ): + await service._send_and_store_message_async( + conversation_id="test-id", + target=target, + request=request, + sequence=0, + request_converter_configurations=[], + response_converter_configurations=[], + preconverted_indexes=set(), + applied_converter_identifiers={}, + ) + + send.assert_not_awaited() + + def _submission( *, conversation_id: str = "main", submission_id: str = "submission", count: int = 1 ) -> MessageSendRequest: From ef07c688f5aaa08b1f1c08e15d65cfe99758bac9 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 9 Oct 2026 12:38:29 -0700 Subject: [PATCH 2/2] Refine shared media validation and fix editor test storage Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 59fd033b-89ce-4424-b418-510b8144b8c8 --- pyrit/backend/routes/media.py | 20 +-- pyrit/backend/services/media_persistence.py | 48 ++++-- .../unit/backend/test_conversation_editor.py | 8 +- tests/unit/backend/test_converter_service.py | 69 ++++++++ tests/unit/backend/test_media_persistence.py | 158 +++++++++--------- tests/unit/backend/test_media_route.py | 6 +- 6 files changed, 200 insertions(+), 109 deletions(-) diff --git a/pyrit/backend/routes/media.py b/pyrit/backend/routes/media.py index 0d6b46ffbd..9ed5f8a79b 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -19,13 +19,11 @@ import asyncio import logging import mimetypes -from pathlib import Path from fastapi import APIRouter, HTTPException, Query from fastapi.responses import FileResponse -from pyrit.backend.services.media_persistence import validate_media_path -from pyrit.memory import CentralMemory +from pyrit.backend.services.media_persistence import MediaAccessDeniedError, validate_local_media_path_async logger = logging.getLogger(__name__) @@ -84,20 +82,14 @@ 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 = await asyncio.to_thread(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 - - try: - validated_path = await asyncio.to_thread(validate_media_path, path=path, allowed_root=allowed_root) - except ValueError as exc: + 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 await asyncio.to_thread(validated_path.is_file): raise HTTPException(status_code=404, detail="File not found.") diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index 66deb83a05..ea2160eac8 100644 --- a/pyrit/backend/services/media_persistence.py +++ b/pyrit/backend/services/media_persistence.py @@ -25,6 +25,13 @@ 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.""" @@ -54,30 +61,49 @@ 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: - ValueError: If the path is outside the allowed media directories. + 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 ValueError("Access denied: path is outside the allowed results directory.") from exc + raise MediaAccessDeniedError("Access denied: path is outside the allowed results directory.") from exc - if not relative_parts or relative_parts[0] not in {"prompt-memory-entries", "seed-prompt-entries"}: - raise ValueError("Access denied: path is not in a media subdirectory.") + 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) -> str: - memory = CentralMemory.get_memory_instance() +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 ValueError("Memory results_path is not configured.") - validated_path = await asyncio.to_thread(validate_media_path, path=path, allowed_root=Path(memory.results_path)) - return str(validated_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: @@ -154,7 +180,7 @@ async def persist_media_value_async( if file_path is None: raise ValueError("Media reference must include a path.") return MediaPersistenceResult( - value=await _validate_local_media_path_async(file_path), + value=str(await validate_local_media_path_async(path=file_path)), origin=MediaOrigin.MEDIA_REFERENCE, persisted=False, resolved=True, @@ -176,7 +202,7 @@ async def persist_media_value_async( is_file = False if is_file: return MediaPersistenceResult( - value=await _validate_local_media_path_async(value), + value=str(await validate_local_media_path_async(path=value)), origin=MediaOrigin.LOCAL_PATH, persisted=False, resolved=True, diff --git a/tests/unit/backend/test_conversation_editor.py b/tests/unit/backend/test_conversation_editor.py index 68732c6a24..56ce1312d3 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, + managed_media_path: Path, data_type: PromptDataType, related: bool, ) -> None: - media = tmp_path / "history.bin" + media = managed_media_path.with_name("history.bin") 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, + managed_media_path: Path, ) -> None: - media = tmp_path / "original.wav" + media = managed_media_path.with_name("original.wav") 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 65560cb681..1d45297cc1 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -15,9 +15,11 @@ import pytest from fastapi import HTTPException +from fastapi.testclient import TestClient from pydantic import ValidationError from pyrit import converter +from pyrit.backend.main import app from pyrit.backend.models.converters import ( ConverterPreviewRequest, CreateConverterRequest, @@ -31,6 +33,8 @@ Base64Converter, BinaryConverter, CaesarConverter, + ImageCompressionConverter, + QRCodeConverter, RepeatTokenConverter, ROT13Converter, SelectiveTextConverter, @@ -1312,6 +1316,44 @@ async def test_preview_conversion_unmarked_media_retains_result_type_async( assert result.steps[0].input_data_type == "image_path" assert result.steps[0].output_data_type == "audio_path" + async def test_preview_media_steps_can_be_served_and_reused_async( + self, *, upload_service: ConverterService + ) -> None: + upload_service._registry.instances.register(QRCodeConverter(), name="qr") + upload_service._registry.instances.register( + ImageCompressionConverter(output_format="PNG", min_compression_threshold=0, fallback_to_original=False), + name="compress", + ) + result = await upload_service.preview_conversion_async( + request=ConverterPreviewRequest(original_value="media preview", converter_ids=["qr", "compress"]) + ) + + assert len(result.steps) == 2 + assert result.steps[0].input_value == "media preview" + assert result.steps[1].input_value == result.steps[0].output_value + assert result.steps[0].output_value != result.steps[1].output_value + assert [step.output_data_type for step in result.steps] == ["image_path", "image_path"] + assert result.converted_value == result.steps[1].output_value + + client = TestClient(app) + for step in result.steps: + response = await asyncio.to_thread(client.get, "/api/media", params={"path": step.output_value}) + assert response.status_code == 200 + assert response.headers["content-type"] == "image/png" + assert response.content == await asyncio.to_thread(Path(step.output_value).read_bytes) + + resumed = await upload_service.preview_conversion_async( + request=ConverterPreviewRequest( + original_value=result.steps[0].output_value, + original_value_data_type=result.steps[0].output_data_type, + converter_ids=["compress"], + ) + ) + assert resumed.steps[0].input_value == result.steps[0].output_value + assert resumed.converted_value_data_type == "image_path" + response = await asyncio.to_thread(client.get, "/api/media", params={"path": resumed.converted_value}) + assert response.status_code == 200 + async def test_preview_conversion_persists_data_uri_for_image_path(self) -> None: """Data URIs on *_path types are decoded via the DEFAULT_MEDIA_EXTENSIONS map and persisted.""" service = ConverterService() @@ -1486,6 +1528,33 @@ async def test_preview_rejects_unmanaged_file_before_conversion_async( assert "outside the allowed results directory" in exc_info.value.detail convert.assert_not_awaited() + @pytest.mark.parametrize("failure", ["uninitialized-memory", "missing-results-path"]) + async def test_preview_storage_configuration_error_returns_500_async( + self, *, upload_service: ConverterService, managed_media_path: Path, failure: str + ) -> None: + memory = CentralMemory.get_memory_instance() + instance = Base64Converter() + upload_service._registry.instances.register(instance, name="base64") + request = ConverterPreviewRequest( + original_value=str(managed_media_path), original_value_data_type="image_path", converter_ids=["base64"] + ) + with ( + patch.object(memory, "results_path", None), + patch.object( + CentralMemory, + "get_memory_instance", + return_value=memory, + side_effect=ValueError("not initialized") if failure == "uninitialized-memory" else None, + ), + patch.object(instance, "convert_async", new_callable=AsyncMock) as convert, + patch.object(converter_routes, "get_converter_service", return_value=upload_service), + pytest.raises(HTTPException) as exc_info, + ): + await converter_routes.preview_conversion(request) + + assert exc_info.value.status_code == 500 + convert.assert_not_awaited() + class TestGetConverterObjectsForIds: """Tests for ConverterService.get_converter_objects_for_ids method.""" diff --git a/tests/unit/backend/test_media_persistence.py b/tests/unit/backend/test_media_persistence.py index e4c597d1a3..047fa072c0 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -11,11 +11,13 @@ from pyrit.backend.models.attacks import MessagePieceRequest from pyrit.backend.services.media_persistence import ( + MediaAccessDeniedError, MediaOrigin, persist_media_value_async, persist_message_pieces_async, + validate_media_path, ) -from pyrit.memory import SQLiteMemory +from pyrit.memory import CentralMemory, SQLiteMemory def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: @@ -56,16 +58,65 @@ async def test_media_reference_requires_path_async(value: str) -> None: factory.assert_not_called() -async def test_existing_local_path_is_not_persisted_async(managed_media_path: Path) -> None: - media_path = managed_media_path - factory = MagicMock() +class TestValidateMediaPath: + @pytest.mark.parametrize("folder", ["prompt-memory-entries", "seed-prompt-entries"]) + def test_accepts_allowed_media_folders(self, *, tmp_path: Path, folder: str) -> None: + path = tmp_path / folder / "nested" / "image.png" - result = await persist_media_value_async(value=str(media_path), data_type="audio_path", serializer_factory=factory) + assert validate_media_path(path=str(path), allowed_root=tmp_path) == path.resolve() - assert result.origin is MediaOrigin.LOCAL_PATH - assert result.value == str(media_path) - assert result.persisted is False - factory.assert_not_called() + @pytest.mark.parametrize("location", ["outside", "results-root", "other-directory", "sibling", "traversal"]) + def test_rejects_unmanaged_paths(self, *, tmp_path: Path, location: str) -> None: + root = tmp_path / "results" + paths = { + "outside": tmp_path / "outside.png", + "results-root": root / "image.png", + "other-directory": root / "other" / "image.png", + "sibling": root / "prompt-memory-entries-other" / "image.png", + "traversal": root / "prompt-memory-entries" / ".." / "image.png", + } + + with pytest.raises(MediaAccessDeniedError, match="Access denied"): + validate_media_path(path=str(paths[location]), allowed_root=root) + + def test_rejects_results_directory_itself(self, *, tmp_path: Path) -> None: + with pytest.raises(MediaAccessDeniedError, match="not in a media subdirectory"): + validate_media_path(path=str(tmp_path), allowed_root=tmp_path) + + def test_returns_canonical_path(self, *, tmp_path: Path) -> None: + path = tmp_path / "prompt-memory-entries" / "nested" / ".." / "image.png" + root = tmp_path / "nested" / ".." + + assert validate_media_path(path=str(path), allowed_root=root) == path.resolve() + + def test_checks_resolved_path_not_input_path(self, *, tmp_path: Path) -> None: + root = tmp_path / "results" + path = root / "prompt-memory-entries" / "image.png" + outside = tmp_path / "outside.png" + original_resolve = Path.resolve + + def resolve(candidate: Path, *, strict: bool = False) -> Path: + return outside if candidate == path else original_resolve(candidate, strict=strict) + + with ( + patch.object(Path, "resolve", new=resolve), + pytest.raises(MediaAccessDeniedError, match="outside the allowed results directory"), + ): + validate_media_path(path=str(path), allowed_root=root) + + def test_rejects_symlink_escape(self, *, tmp_path: Path) -> None: + outside = tmp_path / "outside.png" + outside.write_bytes(b"image") + root = tmp_path / "results" + path = root / "prompt-memory-entries" / "link.png" + path.parent.mkdir(parents=True) + try: + path.symlink_to(outside) + except (OSError, NotImplementedError) as exc: + pytest.skip(f"Cannot create symlink in this environment: {exc}") + + with pytest.raises(MediaAccessDeniedError, match="outside the allowed results directory"): + validate_media_path(path=str(path), allowed_root=root) @pytest.mark.usefixtures("patch_central_database") @@ -73,22 +124,17 @@ class TestLocalMediaPaths: async def test_media_prefix_does_not_bypass_local_path_check_async(self) -> None: with ( patch.object(Path, "is_file", return_value=True), - pytest.raises(ValueError, match="outside the allowed results directory"), + pytest.raises(MediaAccessDeniedError, match="outside the allowed results directory"), ): await persist_media_value_async(value="/api/media-other.png", data_type="image_path") @pytest.mark.parametrize("reference", [False, True]) - @pytest.mark.parametrize("folder", ["prompt-memory-entries", "seed-prompt-entries"]) - async def test_accepts_allowed_media_folders_async( - self, *, managed_media_path: Path, reference: bool, folder: str - ) -> None: - path = managed_media_path.parent.parent / folder / "nested" / "image.png" - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(b"image") + async def test_existing_media_is_not_persisted_async(self, *, managed_media_path: Path, reference: bool) -> None: + path = managed_media_path value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) factory = MagicMock() - result = await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) + result = await persist_media_value_async(value=value, data_type="audio_path", serializer_factory=factory) assert result.value == str(path.resolve()) assert result.origin is (MediaOrigin.MEDIA_REFERENCE if reference else MediaOrigin.LOCAL_PATH) @@ -97,68 +143,17 @@ async def test_accepts_allowed_media_folders_async( factory.assert_not_called() @pytest.mark.parametrize("reference", [False, True]) - @pytest.mark.parametrize("location", ["outside", "results-root", "other-directory", "sibling", "traversal"]) - async def test_rejects_unmanaged_files_async( - self, *, managed_media_path: Path, tmp_path: Path, reference: bool, location: str - ) -> None: - root = managed_media_path.parent.parent - paths = { - "outside": tmp_path / "outside.png", - "results-root": root / "image.png", - "other-directory": root / "other" / "image.png", - "sibling": root / "prompt-memory-entries-other" / "image.png", - "traversal": managed_media_path.parent / ".." / "image.png", - } - path = paths[location] - path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(b"image") - value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) + async def test_rejects_unmanaged_file_without_persistence_async(self, *, tmp_path: Path, reference: bool) -> None: + outside = tmp_path / "outside.png" + outside.write_bytes(b"image") + value = f"/api/media?{urlencode({'path': str(outside)})}" if reference else str(outside) factory = MagicMock() - with pytest.raises(ValueError, match="Access denied"): + with pytest.raises(MediaAccessDeniedError, match="outside the allowed results directory"): await persist_media_value_async(value=value, data_type="image_path", serializer_factory=factory) factory.assert_not_called() - @pytest.mark.parametrize("reference", [False, True]) - async def test_returns_canonical_allowed_path_async(self, *, managed_media_path: Path, reference: bool) -> None: - nested = managed_media_path.parent / "nested" - nested.mkdir() - path = nested / ".." / managed_media_path.name - value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) - - result = await persist_media_value_async(value=value, data_type="image_path") - - assert result.value == str(managed_media_path.resolve()) - - @pytest.mark.parametrize("reference", [False, True]) - async def test_checks_resolved_path_not_input_path_async( - self, *, managed_media_path: Path, tmp_path: Path, reference: bool - ) -> None: - outside = tmp_path / "outside.png" - value = f"/api/media?{urlencode({'path': str(managed_media_path)})}" if reference else str(managed_media_path) - with ( - patch.object(Path, "resolve", side_effect=[outside, managed_media_path.parent.parent]), - pytest.raises(ValueError, match="outside the allowed results directory"), - ): - await persist_media_value_async(value=value, data_type="image_path") - - @pytest.mark.parametrize("reference", [False, True]) - async def test_rejects_symlink_escape_async( - self, *, managed_media_path: Path, tmp_path: Path, reference: bool - ) -> None: - outside = tmp_path / "outside.png" - outside.write_bytes(b"image") - path = managed_media_path.parent / "link.png" - try: - path.symlink_to(outside) - except (OSError, NotImplementedError) as exc: - pytest.skip(f"Cannot create symlink in this environment: {exc}") - value = f"/api/media?{urlencode({'path': str(path)})}" if reference else str(path) - - with pytest.raises(ValueError, match="outside the allowed results directory"): - await persist_media_value_async(value=value, data_type="image_path") - @pytest.mark.parametrize("reference", [False, True]) @pytest.mark.parametrize("field", ["original_value", "converted_value"]) async def test_checks_both_message_values_async( @@ -176,18 +171,27 @@ async def test_checks_both_message_values_async( before = piece.model_dump() factory = MagicMock() - with pytest.raises(ValueError, match="outside the allowed results directory"): + with pytest.raises(MediaAccessDeniedError, match="outside the allowed results directory"): await persist_message_pieces_async(pieces=[piece], serializer_factory=factory) assert piece.model_dump() == before factory.assert_not_called() + @pytest.mark.parametrize("reference", [False, True]) async def test_requires_configured_results_path_async( - self, *, managed_media_path: Path, sqlite_instance: SQLiteMemory + self, *, managed_media_path: Path, sqlite_instance: SQLiteMemory, reference: bool ) -> None: + value = f"/api/media?{urlencode({'path': str(managed_media_path)})}" if reference else str(managed_media_path) with ( patch.object(sqlite_instance, "results_path", None), - pytest.raises(ValueError, match="results_path is not configured"), + pytest.raises(RuntimeError, match="results_path is not configured"), + ): + await persist_media_value_async(value=value, data_type="image_path") + + async def test_requires_initialized_memory_async(self, *, managed_media_path: Path) -> None: + with ( + patch.object(CentralMemory, "get_memory_instance", side_effect=ValueError("not initialized")), + pytest.raises(RuntimeError, match="Memory not initialized"), ): await persist_media_value_async(value=str(managed_media_path), data_type="image_path") diff --git a/tests/unit/backend/test_media_route.py b/tests/unit/backend/test_media_route.py index efb52ede86..0e8a4434e8 100644 --- a/tests/unit/backend/test_media_route.py +++ b/tests/unit/backend/test_media_route.py @@ -30,7 +30,7 @@ def _mock_memory(tmp_path: Path): # Create allowed subdirectories (tmp_path / "prompt-memory-entries").mkdir() (tmp_path / "seed-prompt-entries").mkdir() - with patch("pyrit.backend.routes.media.CentralMemory") as mock_cm: + with patch("pyrit.backend.services.media_persistence.CentralMemory") as mock_cm: mock_cm.get_memory_instance.return_value = mock_mem yield tmp_path @@ -196,7 +196,7 @@ class TestServeMediaErrors: def test_returns_500_when_memory_not_initialized(self, client: TestClient) -> None: """Returns 500 when CentralMemory is not initialized.""" - with patch("pyrit.backend.routes.media.CentralMemory") as mock_cm: + with patch("pyrit.backend.services.media_persistence.CentralMemory") as mock_cm: mock_cm.get_memory_instance.side_effect = ValueError("not initialized") response = client.get("/api/media", params={"path": "/some/file.png"}) @@ -207,7 +207,7 @@ def test_returns_500_when_results_path_is_none(self, client: TestClient) -> None """Returns 500 when memory.results_path is None.""" mock_mem = MagicMock() mock_mem.results_path = None - with patch("pyrit.backend.routes.media.CentralMemory") as mock_cm: + with patch("pyrit.backend.services.media_persistence.CentralMemory") as mock_cm: mock_cm.get_memory_instance.return_value = mock_mem response = client.get("/api/media", params={"path": "/some/file.png"})