diff --git a/pyrit/backend/routes/media.py b/pyrit/backend/routes/media.py index da93d1e3e1..9ed5f8a79b 100644 --- a/pyrit/backend/routes/media.py +++ b/pyrit/backend/routes/media.py @@ -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 = { @@ -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."), @@ -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() diff --git a/pyrit/backend/services/media_persistence.py b/pyrit/backend/services/media_persistence.py index 4985697c05..ea2160eac8 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: @@ -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.""" @@ -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: @@ -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. """ @@ -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 @@ -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, 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_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 ed39ed57d7..1d45297cc1 100644 --- a/tests/unit/backend/test_converter_service.py +++ b/tests/unit/backend/test_converter_service.py @@ -11,12 +11,15 @@ 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 +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, @@ -30,6 +33,8 @@ Base64Converter, BinaryConverter, CaesarConverter, + ImageCompressionConverter, + QRCodeConverter, RepeatTokenConverter, ROT13Converter, SelectiveTextConverter, @@ -1019,14 +1024,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 +1296,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,14 +1308,52 @@ 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" + 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() @@ -1442,11 +1489,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 +1505,56 @@ 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() + + @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 2fdab2f1af..047fa072c0 100644 --- a/tests/unit/backend/test_media_persistence.py +++ b/tests/unit/backend/test_media_persistence.py @@ -5,10 +5,19 @@ 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 ( + MediaAccessDeniedError, + MediaOrigin, + persist_media_value_async, + persist_message_pieces_async, + validate_media_path, +) +from pyrit.memory import CentralMemory, SQLiteMemory def _serializer(*, value: str = "/saved/media.bin") -> MagicMock: @@ -22,8 +31,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,19 +48,154 @@ 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() - result = await persist_media_value_async(value=str(media_path), data_type="audio_path", serializer_factory=factory) + 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) - assert result.origin is MediaOrigin.LOCAL_PATH - assert result.value == str(media_path) - assert result.persisted is False factory.assert_not_called() +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" + + assert validate_media_path(path=str(path), allowed_root=tmp_path) == path.resolve() + + @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") +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(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]) + 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="audio_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]) + 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(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]) + @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(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, 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(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") + + 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_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"}) 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: