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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions pyrit/memory/storage/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,12 @@ async def save_formatted_audio_async(
Raises:
RuntimeError: If storage IO is not initialized.
"""
original_file_extension = self.file_extension
self.file_extension = "wav"
if output_filename:
output_suffix = Path(output_filename).suffix
if output_suffix.casefold() == f".{original_file_extension}".casefold():
output_filename = output_filename[: -len(output_suffix)]
file_path = await self.get_data_filename_async(file_name=output_filename)

# save audio file locally first if in AzureStorageBlob so we can use wave.open to set audio parameters
Expand Down Expand Up @@ -346,6 +352,9 @@ async def get_data_filename_async(self, file_name: str | None = None) -> Path |

results_path = str(DB_DATA_PATH)
file_name = file_name if file_name else str(ticks)
file_suffix = Path(file_name).suffix
Comment thread
romanlutz marked this conversation as resolved.
if file_suffix.casefold() == f".{self.file_extension}".casefold():
file_name = file_name[: -len(file_suffix)]

if self._is_azure_storage_url(results_path):
full_data_directory_path = results_path + self.data_sub_directory
Expand Down
41 changes: 39 additions & 2 deletions pyrit/memory/storage/storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from abc import ABC, abstractmethod
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, ClassVar
from urllib.parse import urlparse

import aiofiles
Expand Down Expand Up @@ -157,6 +157,42 @@ class AzureBlobStorageIO(StorageIO):
Implementation of StorageIO for Azure Blob Storage.
"""

_EXTENSION_TO_CONTENT_TYPE: ClassVar[dict[str, str]] = {
Comment thread
romanlutz marked this conversation as resolved.
".txt": "text/plain",
".html": "text/html",
".htm": "text/html",
".csv": "text/csv",
".md": "text/markdown",
".json": "application/json",
".xml": "application/xml",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".gif": "image/gif",
".webp": "image/webp",
".svg": "image/svg+xml",
".bmp": "image/bmp",
".wav": "audio/wav",
".mp3": "audio/mpeg",
".ogg": "audio/ogg",
".flac": "audio/flac",
".m4a": "audio/mp4",
".mp4": "video/mp4",
".webm": "video/webm",
".ogv": "video/ogg",
".avi": "video/x-msvideo",
".pdf": "application/pdf",
".doc": "application/msword",
".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
".xls": "application/vnd.ms-excel",
".xlsx": "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
".ppt": "application/vnd.ms-powerpoint",
".pptx": "application/vnd.openxmlformats-officedocument.presentationml.presentation",
".rtf": "application/rtf",
".zip": "application/zip",
".bin": "application/octet-stream",
}

def __init__(
self,
*,
Expand Down Expand Up @@ -382,8 +418,9 @@ async def write_file_async(self, path: Path | str, data: bytes) -> None:
if not self._client_async:
self._client_async = await self._create_container_client_async()
blob_name = self._resolve_blob_name(path)
content_type = self._EXTENSION_TO_CONTENT_TYPE.get(Path(blob_name).suffix.lower(), self._blob_content_type)
try:
await self._upload_blob_async(file_name=blob_name, data=data, content_type=self._blob_content_type)
await self._upload_blob_async(file_name=blob_name, data=data, content_type=content_type)
except Exception as exc:
logger.exception(f"Failed to write file at {blob_name}: {exc}")
raise
Expand Down
114 changes: 114 additions & 0 deletions tests/unit/memory/storage/test_serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import os
import re
import tempfile
from pathlib import Path
from typing import get_args
from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch

Expand All @@ -13,10 +14,12 @@

from pyrit.memory.storage import (
AllowedCategories,
AzureBlobStorageIO,
BinaryPathDataTypeSerializer,
DataTypeSerializer,
ErrorDataTypeSerializer,
ImagePathDataTypeSerializer,
StorageIO,
TextDataTypeSerializer,
data_serializer_factory,
set_message_piece_sha256_async,
Expand All @@ -25,6 +28,35 @@
from pyrit.models import MessagePiece, SeedPrompt


class LegacyStorageIO(StorageIO):
"""
Test double representing an existing third-party ``StorageIO`` implementation.

Its ``write_file_async(path, data)`` method intentionally retains the original
two-argument contract. Tests using this class ensure serializers do not pass a
new content-type keyword argument that would break pre-existing custom storage
backends when Azure Blob Storage adds MIME metadata internally.
"""

def __init__(self) -> None:
self.writes: list[tuple[Path | str, bytes]] = []

async def read_file_async(self, path: Path | str) -> bytes:
return b""

async def write_file_async(self, path: Path | str, data: bytes) -> None:
self.writes.append((path, data))

async def path_exists_async(self, path: Path | str) -> bool:
return False

async def is_file_async(self, path: Path | str) -> bool:
return False

async def create_directory_if_not_exists_async(self, path: Path | str) -> None:
return None


def test_allowed_categories():
entries = get_args(AllowedCategories)
assert len(entries) == 2
Expand Down Expand Up @@ -285,6 +317,41 @@ async def test_get_data_filename(sqlite_instance):
assert not os.path.exists(filename) # File should not exist yet


async def test_get_data_filename_does_not_duplicate_extension(sqlite_instance):
serializer = data_serializer_factory(category="prompt-memory-entries", data_type="image_path")

filename = await serializer.get_data_filename_async(file_name="photo.png")

assert Path(filename).name == "photo.png"


async def test_get_data_filename_preserves_dotted_basename(sqlite_instance):
serializer = data_serializer_factory(
category="prompt-memory-entries",
data_type="binary_path",
extension="pdf",
)

filename = await serializer.get_data_filename_async(file_name="2024.10.15_report")

assert Path(filename).name == "2024.10.15_report.pdf"


async def test_save_data_supports_legacy_storage_io_write_signature():
storage = LegacyStorageIO()
mock_memory = MagicMock()
mock_memory.results_path = "https://account.blob.core.windows.net/container/results"
mock_memory.results_storage_io = storage
serializer = data_serializer_factory(category="prompt-memory-entries", data_type="image_path")

with patch.object(type(serializer), "_memory", new_callable=PropertyMock, return_value=mock_memory):
await serializer.save_data_async(b"\x89PNG", output_filename="photo.png")

assert storage.writes == [
("https://account.blob.core.windows.net/container/results/prompt-memory-entries/images/photo.png", b"\x89PNG")
]


def test_binary_path_normalizer_factory(sqlite_instance):
"""Test factory creates BinaryPathDataTypeSerializer correctly."""
serializer = data_serializer_factory(category="prompt-memory-entries", data_type="binary_path")
Expand Down Expand Up @@ -317,6 +384,23 @@ async def test_binary_path_save_data(sqlite_instance):
assert os.path.isfile(serializer_value)


async def test_binary_path_default_extension_sets_azure_content_type():
storage = AzureBlobStorageIO(container_url="https://account.blob.core.windows.net/container")
mock_container_client = AsyncMock()
storage._client_async = mock_container_client
mock_memory = MagicMock()
mock_memory.results_path = "https://account.blob.core.windows.net/container/results"
mock_memory.results_storage_io = storage
serializer = data_serializer_factory(category="prompt-memory-entries", data_type="binary_path")

with patch.object(type(serializer), "_memory", new_callable=PropertyMock, return_value=mock_memory):
await serializer.save_data_async(b"\x00\x01", output_filename="payload")

upload_kwargs = mock_container_client.upload_blob.await_args.kwargs
assert upload_kwargs["name"] == "results/prompt-memory-entries/binaries/payload.bin"
assert upload_kwargs["content_settings"].content_type == "application/octet-stream"


async def test_binary_path_read_data(sqlite_instance):
"""Test reading binary data from disk."""
data = b"\x00\x11\x22\x33\x44\x55"
Expand Down Expand Up @@ -433,6 +517,36 @@ async def test_save_formatted_audio_writes_local_wav_via_to_thread(sqlite_instan
assert wav_file.readframes(wav_file.getnframes()) == pcm


async def test_save_formatted_audio_uses_wav_filename_content_and_metadata(tmp_path):
import io
import wave

storage = AzureBlobStorageIO(container_url="https://account.blob.core.windows.net/container")
mock_container_client = AsyncMock()
storage._client_async = mock_container_client
mock_memory = MagicMock()
mock_memory.results_path = "https://account.blob.core.windows.net/container/results"
mock_memory.results_storage_io = storage
serializer = data_serializer_factory(category="prompt-memory-entries", data_type="audio_path")
pcm = b"\x01\x00\x02\x00\x03\x00\x04\x00"

with (
patch.object(type(serializer), "_memory", new_callable=PropertyMock, return_value=mock_memory),
patch("pyrit.memory.storage.serializers.DB_DATA_PATH", tmp_path),
):
await serializer.save_formatted_audio_async(data=pcm, output_filename="recording.mp3")

upload_kwargs = mock_container_client.upload_blob.await_args.kwargs
assert upload_kwargs["name"] == "results/prompt-memory-entries/audio/recording.wav"
assert upload_kwargs["content_settings"].content_type == "audio/wav"
assert serializer.value.endswith("/audio/recording.wav")
with wave.open(io.BytesIO(upload_kwargs["data"]), "rb") as wav_file:
assert wav_file.getnchannels() == 1
assert wav_file.getsampwidth() == 2
assert wav_file.getframerate() == 16000
assert wav_file.readframes(wav_file.getnframes()) == pcm


def test_write_wav_sync_produces_readable_wav(tmp_path):
"""_write_wav_sync should produce a WAV file readable by wave.open with the same metadata and frames."""
import wave
Expand Down
37 changes: 37 additions & 0 deletions tests/unit/memory/storage/test_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,6 +162,43 @@ async def test_azure_blob_storage_io_write_file_with_relative_path():
)


@pytest.mark.parametrize(
("path", "expected_content_type"),
[
("notes.HTML", "text/html"),
("photo.JPEG", "image/jpeg"),
("recording.wav", "audio/wav"),
("movie.mp4", "video/mp4"),
("report.pdf", "application/pdf"),
],
)
async def test_azure_blob_storage_io_write_file_sets_content_type_from_extension(path, expected_content_type):
storage = AzureBlobStorageIO(container_url="https://account.blob.core.windows.net/container")
mock_container_client = AsyncMock()
storage._client_async = mock_container_client

await storage.write_file_async(path, b"data")

upload_kwargs = mock_container_client.upload_blob.await_args.kwargs
assert upload_kwargs["name"] == path
assert upload_kwargs["content_settings"].content_type == expected_content_type


@pytest.mark.parametrize("path", ["data.unknown", "README"])
async def test_azure_blob_storage_io_write_file_uses_configured_fallback_content_type(path):
storage = AzureBlobStorageIO(
container_url="https://account.blob.core.windows.net/container",
blob_content_type=SupportedContentType.PLAIN_TEXT,
)
mock_container_client = AsyncMock()
storage._client_async = mock_container_client

await storage.write_file_async(path, b"data")

upload_kwargs = mock_container_client.upload_blob.await_args.kwargs
assert upload_kwargs["content_settings"].content_type == SupportedContentType.PLAIN_TEXT.value


async def test_azure_blob_storage_io_create_container_client_uses_explicit_sas_token():
container_url = "https://youraccount.blob.core.windows.net/yourcontainer"
sas_token = "explicit-sas-token"
Expand Down
Loading