diff --git a/pyrit/datasets/seed_datasets/remote/remote_dataset_loader.py b/pyrit/datasets/seed_datasets/remote/remote_dataset_loader.py index 2f17525065..c4866fcecc 100644 --- a/pyrit/datasets/seed_datasets/remote/remote_dataset_loader.py +++ b/pyrit/datasets/seed_datasets/remote/remote_dataset_loader.py @@ -361,13 +361,15 @@ def _load_dataset_sync() -> Any: """ cache_dir = str(DB_DATA_PATH / "huggingface") if cache else None - # Explicitly set download_mode to reuse cached data and never re-download + # Reuse cached data when caching is enabled; force a re-download otherwise so + # cache=False genuinely picks up upstream edits instead of silently reusing the cache. + download_mode = DownloadMode.REUSE_DATASET_IF_EXISTS if cache else DownloadMode.FORCE_REDOWNLOAD return load_dataset( dataset_name, config, split=split, cache_dir=cache_dir, - download_mode=DownloadMode.REUSE_DATASET_IF_EXISTS, + download_mode=download_mode, token=token, **kwargs, ) diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 602f13b6bf..4e46b73b5f 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -2205,6 +2205,44 @@ async def _serialize_seed_value_async(self, prompt: Seed) -> str: serialized_prompt_value = str(serializer.value) return serialized_prompt_value or "" + async def _prepare_seed_for_storage_async( + self, *, prompt: Seed, added_by: str | None, current_time: datetime + ) -> None: + """ + Prepare a seed in place for persistence. + + Sets provenance and timestamp, serializes any media value to storage, and computes the + SHA256 used for identity and deduplication. Performs no database writes, so it is safe to + call before opening a transaction. + + Args: + prompt (Seed): The seed to prepare; it is mutated in place. + added_by (str | None): The user to attribute the seed to; overrides an existing value. + current_time (datetime): The timestamp to apply when the seed has no ``date_added``. + + Raises: + ValueError: If ``added_by`` is not set on the seed and none is provided. + """ + if added_by: + prompt.added_by = added_by + if not prompt.added_by: + raise ValueError( + """The 'added_by' attribute must be set for each prompt. + Set it explicitly or pass a value to the 'added_by' parameter.""" + ) + if prompt.date_added is None: + prompt.date_added = current_time + + # Only SeedPrompt has set_encoding_metadata for audio/video/image files + if hasattr(prompt, "set_encoding_metadata"): + prompt.set_encoding_metadata() # type: ignore[ty:call-non-callable] + + # Handle serialization for image, audio & video SeedPrompts + if prompt.data_type in ["image_path", "audio_path", "video_path"]: + prompt.value = await self._serialize_seed_value_async(prompt=prompt) + + await set_seed_sha256_async(prompt) + async def add_seeds_to_memory_async(self, *, seeds: Sequence[Seed], added_by: str | None = None) -> None: """ Insert a list of seeds into the memory storage. @@ -2219,26 +2257,7 @@ async def add_seeds_to_memory_async(self, *, seeds: Sequence[Seed], added_by: st entries: MutableSequence[SeedEntry] = [] current_time = datetime.now(tz=timezone.utc) for prompt in seeds: - if added_by: - prompt.added_by = added_by - if not prompt.added_by: - raise ValueError( - """The 'added_by' attribute must be set for each prompt. - Set it explicitly or pass a value to the 'added_by' parameter.""" - ) - if prompt.date_added is None: - prompt.date_added = current_time - - # Only SeedPrompt has set_encoding_metadata for audio/video/image files - if hasattr(prompt, "set_encoding_metadata"): - prompt.set_encoding_metadata() # type: ignore[ty:call-non-callable] - - # Handle serialization for image, audio & video SeedPrompts - if prompt.data_type in ["image_path", "audio_path", "video_path"]: - serialized_prompt_value = await self._serialize_seed_value_async(prompt=prompt) - prompt.value = serialized_prompt_value - - await set_seed_sha256_async(prompt) + await self._prepare_seed_for_storage_async(prompt=prompt, added_by=added_by, current_time=current_time) if prompt.value_sha256 and not self.get_seeds( value_sha256=[prompt.value_sha256], dataset_name=prompt.dataset_name @@ -2281,6 +2300,81 @@ def get_seed_dataset_names(self) -> Sequence[str]: logger.exception(f"Failed to retrieve dataset names with error {e}") raise + async def replace_seeds_for_dataset_async( + self, *, dataset_name: str, seeds: Sequence[Seed], added_by: str | None = None + ) -> int: + """ + Atomically replace all stored seeds for a dataset with a new set. + + Every existing ``SeedPromptEntries`` row for ``dataset_name`` is deleted and the provided + seeds are inserted in a single transaction and commit; if the insert fails the delete is + rolled back with it, so the previously stored seeds are preserved. Seeds are prepared + (media serialized, SHA256 computed) before the transaction opens. Deduplication is + intentionally skipped: this is a full replace, so the provided seeds are stored as given. + + The isolation guarantee is the database transaction boundary: a reader that queries after + the commit sees the complete new set. This holds on the file-backed SQLite and Azure SQL + backends, where each session has its own connection. The in-memory SQLite backend shares a + single connection across all sessions, so it does not isolate concurrent sessions from one + another; callers that need to read a dataset while it is being replaced should use a + file-backed or Azure SQL backend. ``RefreshDatasets`` replaces datasets sequentially, so it + does not rely on cross-session isolation. + + ``SeedPromptEntries`` has no dependent foreign keys, so no related rows are removed first. + Deleting media-backed seeds (``image_path``, ``audio_path``, ``video_path``) removes only the + database rows; any serialized media files they reference are left on disk. This matches every + other seed-delete path and results in disk bloat, not data loss. + + Args: + dataset_name (str): The name of the dataset whose seeds should be replaced. + seeds (Sequence[Seed]): The new seeds to store for the dataset; must be non-empty and + every seed's ``dataset_name`` must equal ``dataset_name``. + added_by (str | None): The user to attribute the new seeds to. + + Returns: + int: The number of ``SeedPromptEntries`` deleted before the new seeds were inserted. + + Raises: + ValueError: If ``dataset_name`` is empty, ``seeds`` is empty, or any seed's + ``dataset_name`` does not match ``dataset_name``. + SQLAlchemyError: If the replacement fails; the transaction is rolled back first. + """ + if not dataset_name: + raise ValueError("dataset_name must be a non-empty string.") + if not seeds: + raise ValueError("seeds must be non-empty; refusing to replace a dataset with nothing.") + mismatched = sorted( + {seed.dataset_name for seed in seeds if seed.dataset_name != dataset_name}, + key=lambda name: (name is None, name or ""), + ) + if mismatched: + raise ValueError( + f"All seeds must belong to dataset '{dataset_name}', but got mismatched " + f"dataset_name(s): {mismatched}. Refusing to delete '{dataset_name}' and insert " + "seeds tagged for another dataset." + ) + + current_time = datetime.now(tz=timezone.utc) + entries: list[SeedEntry] = [] + for prompt in seeds: + await self._prepare_seed_for_storage_async(prompt=prompt, added_by=added_by, current_time=current_time) + entries.append(SeedEntry(entry=prompt)) + + with closing(self.get_session()) as session: + try: + deleted = ( + session.query(SeedEntry) + .filter(SeedEntry.dataset_name == dataset_name) + .delete(synchronize_session=False) + ) + session.add_all(entries) + session.commit() + return deleted + except SQLAlchemyError as e: + session.rollback() + logger.exception(f"Error replacing seeds for dataset {dataset_name}: {e}") + raise + async def add_seed_groups_to_memory_async( self, *, prompt_groups: Sequence[SeedGroup], added_by: str | None = None ) -> None: diff --git a/pyrit/setup/initializers/__init__.py b/pyrit/setup/initializers/__init__.py index 1e240e5b63..69cb3d4c11 100644 --- a/pyrit/setup/initializers/__init__.py +++ b/pyrit/setup/initializers/__init__.py @@ -6,6 +6,7 @@ from pyrit.models.parameter import Parameter from pyrit.setup.initializers.load_default_datasets import LoadDefaultDatasets from pyrit.setup.initializers.preload_scenario_metadata import PreloadScenarioMetadata +from pyrit.setup.initializers.refresh_datasets import RefreshDatasets from pyrit.setup.initializers.scorers import ScorerInitializer from pyrit.setup.initializers.targets import TargetInitializer from pyrit.setup.initializers.techniques import TechniqueInitializer @@ -19,4 +20,5 @@ "TargetInitializer", "LoadDefaultDatasets", "PreloadScenarioMetadata", + "RefreshDatasets", ] diff --git a/pyrit/setup/initializers/refresh_datasets.py b/pyrit/setup/initializers/refresh_datasets.py new file mode 100644 index 0000000000..3d4deaa56b --- /dev/null +++ b/pyrit/setup/initializers/refresh_datasets.py @@ -0,0 +1,254 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Refresh datasets already loaded into memory. + +Re-fetches datasets that are present in ``CentralMemory`` from their registered providers and +replaces their stored seeds, so previously loaded copies pick up upstream changes such as +standardized harm categories, corrected metadata, or live threat-feed updates. This is the +maintenance twin of ``LoadDefaultDatasets``: it is opt-in and never runs on the scenario hot path. +""" + +import logging +import textwrap +from datetime import datetime, timedelta, timezone + +from pyrit.datasets import SeedDatasetProvider +from pyrit.memory import CentralMemory, MemoryInterface +from pyrit.models import SeedDataset +from pyrit.models.parameter import Parameter +from pyrit.setup.pyrit_initializer import PyRITInitializer + +logger = logging.getLogger(__name__) + + +class RefreshDatasets(PyRITInitializer): + """ + Refresh datasets already loaded in memory from their registered providers. + + For each selected dataset that is present in memory and backed by a registered provider, this + re-fetches the dataset with caching disabled and replaces its stored seeds. Selection can be + narrowed with ``dataset_names``; a ``days`` threshold limits the refresh to datasets whose + newest seed is older than ``days`` days (``days=0`` refreshes every selected dataset regardless + of age). + + Datasets in memory that have no registered provider (for example custom, manually added ones) + are skipped, since there is nothing to re-fetch. + """ + + DEFAULT_DAYS: int = 30 + ADDED_BY: str = "RefreshDatasets" + + @property + def description(self) -> str: + """A description of this initializer.""" + return textwrap.dedent( + """ + Refreshes datasets already present in memory by re-fetching them from their + registered providers with caching disabled and replacing their stored seeds. Use + days to refresh only datasets whose newest seed is older than N days (days=0 + refreshes all selected datasets); use dataset_names to narrow the selection. + + Note: this is intended for periodic maintenance, not the scenario hot path. It only + refreshes datasets that are already in memory and backed by a registered provider. + """ + ).strip() + + @property + def required_env_vars(self) -> list[str]: + """The list of required environment variables.""" + return [] + + @property + def supported_parameters(self) -> list[Parameter]: + """The list of parameters this initializer accepts.""" + return [ + Parameter( + name="days", + description=( + "Refresh only datasets whose newest seed is older than this many days. " + "0 refreshes every selected dataset regardless of age." + ), + default=self.DEFAULT_DAYS, + ), + Parameter( + name="dataset_names", + description="Explicit dataset names to refresh; refreshes all in-memory datasets if omitted.", + default=[], + ), + ] + + async def initialize_async(self) -> None: + """Refresh the selected stale datasets in CentralMemory, isolating per-dataset failures.""" + days = self._parse_days() + memory = CentralMemory.get_memory_instance() + + names_in_memory = set(memory.get_seed_dataset_names()) + if not names_in_memory: + logger.warning("No datasets in memory to refresh") + return + + candidates = await self._select_candidates_async(names_in_memory=names_in_memory) + if not candidates: + logger.warning("No datasets matched the requested selection") + return + + refreshed: list[str] = [] + up_to_date: list[str] = [] + failed: list[str] = [] + for name in candidates: + if not self._is_stale(memory=memory, dataset_name=name, days=days): + up_to_date.append(name) + continue + try: + await self._refresh_dataset_async(memory=memory, dataset_name=name) + refreshed.append(name) + except Exception as exc: # noqa: BLE001 - isolate one dataset's failure from the rest + logger.warning(f"Skipping refresh for dataset '{name}': {exc}") + failed.append(name) + + logger.info(f"Refresh complete: {len(refreshed)} refreshed, {len(up_to_date)} up-to-date, {len(failed)} failed") + + async def _select_candidates_async(self, *, names_in_memory: set[str]) -> list[str]: + """ + Resolve which in-memory datasets to consider for refresh. + + With explicit ``dataset_names``, only those are considered; otherwise every in-memory + dataset that has a registered provider is considered. Names that are not in memory or have + no registered provider are skipped with a log message. + + Args: + names_in_memory (set[str]): The dataset names currently present in memory. + + Returns: + list[str]: The dataset names to evaluate for staleness. + """ + dataset_names = self.params.get("dataset_names", []) + registered = set(await SeedDatasetProvider.get_all_dataset_names_async()) + + if dataset_names: + candidates: list[str] = [] + for name in dict.fromkeys(dataset_names): + if name not in names_in_memory: + logger.warning(f"Skipping '{name}': not present in memory") + elif name not in registered: + logger.warning(f"Skipping '{name}': no registered provider to refresh from") + else: + candidates.append(name) + return candidates + + selected: list[str] = [] + for name in sorted(names_in_memory): + if name in registered: + selected.append(name) + else: + logger.debug(f"Skipping '{name}': no registered provider to refresh from") + return selected + + def _is_stale(self, *, memory: MemoryInterface, dataset_name: str, days: int) -> bool: + """ + Determine whether a dataset is stale enough to refresh. + + Args: + memory (MemoryInterface): The memory instance to read existing seeds from. + dataset_name (str): The dataset to evaluate. + days (int): The staleness threshold in days; 0 always refreshes. + + Returns: + bool: True if the dataset should be refreshed, otherwise False. + """ + if days == 0: + return True + + seeds = memory.get_seeds(dataset_name=dataset_name) + newest = max((seed.date_added for seed in seeds if seed.date_added is not None), default=None) + if newest is None: + return True + + cutoff = datetime.now(tz=timezone.utc) - timedelta(days=days) + return newest <= cutoff + + async def _refresh_dataset_async(self, *, memory: MemoryInterface, dataset_name: str) -> None: + """ + Re-fetch a single dataset and atomically replace its stored seeds. + + The dataset is fetched (with caching disabled) before anything is deleted, and the replace + is a single transaction, so a failed or empty fetch - or a failed insert - leaves the + existing seeds untouched. + + Args: + memory (MemoryInterface): The memory instance to replace seeds in. + dataset_name (str): The dataset to refresh. + + Raises: + ValueError: If the provider returns no usable dataset for ``dataset_name``. + """ + fetched = await SeedDatasetProvider.fetch_datasets_async( + dataset_names=[dataset_name], cache=False, max_concurrency=1 + ) + dataset = self._require_matching_dataset(fetched=fetched, dataset_name=dataset_name) + + deleted = await memory.replace_seeds_for_dataset_async( + dataset_name=dataset_name, seeds=dataset.seeds, added_by=self.ADDED_BY + ) + logger.info(f"Refreshed dataset '{dataset_name}': replaced {deleted} seeds with {len(dataset.seeds)}") + + def _parse_days(self) -> int: + """ + Parse and validate the ``days`` parameter. + + Returns: + int: The validated non-negative staleness threshold. + + Raises: + ValueError: If ``days`` is not a single non-negative integer. + """ + raw = self.params.get("days", []) + if not raw: + return self.DEFAULT_DAYS + if len(raw) != 1: + raise ValueError(f"'days' must be a single non-negative integer, got {raw}") + try: + days = int(raw[0]) + except (TypeError, ValueError): + raise ValueError(f"'days' must be a non-negative integer, got {raw[0]!r}") from None + if days < 0: + raise ValueError(f"'days' must be non-negative, got {days}") + return days + + @staticmethod + def _require_matching_dataset(*, fetched: list[SeedDataset], dataset_name: str) -> SeedDataset: + """ + Validate a fetch returned exactly the requested, non-empty dataset before replacing seeds. + + Guards the destructive replace against a provider that returns nothing, more than one + dataset, an empty dataset, or a dataset whose seeds carry a different ``dataset_name`` than + requested (which would delete the requested dataset and insert unrelated seeds). + + Args: + fetched (list[SeedDataset]): The datasets returned by the provider. + dataset_name (str): The dataset name that was requested. + + Returns: + SeedDataset: The single fetched dataset that matches ``dataset_name``. + + Raises: + ValueError: If the fetch did not return exactly one non-empty dataset for the + requested name. + """ + if len(fetched) != 1: + raise ValueError(f"Expected exactly one dataset for '{dataset_name}', got {len(fetched)}") + dataset = fetched[0] + if not dataset.seeds: + raise ValueError(f"Re-fetched dataset '{dataset_name}' is empty; keeping existing seeds") + mismatched = sorted( + {seed.dataset_name for seed in dataset.seeds if seed.dataset_name != dataset_name}, + key=lambda name: (name is None, name or ""), + ) + if mismatched: + raise ValueError( + f"Re-fetched dataset for '{dataset_name}' contains seeds for other datasets " + f"{mismatched}; keeping existing seeds" + ) + return dataset diff --git a/tests/unit/datasets/test_remote_dataset_loader.py b/tests/unit/datasets/test_remote_dataset_loader.py index a9a274fa56..1978e48425 100644 --- a/tests/unit/datasets/test_remote_dataset_loader.py +++ b/tests/unit/datasets/test_remote_dataset_loader.py @@ -291,3 +291,27 @@ async def test_unsupported_inner_extension_raises_valueerror(self): loader = ConcreteRemoteLoader() with pytest.raises(ValueError, match="Invalid file_type"): await loader._fetch_zip_from_url_async(source=self.SOURCE, inner_files=["bad.parquet"], cache=False) + + +class TestFetchFromHuggingFaceDownloadMode: + """The cache flag must drive the HuggingFace download_mode so cache=False re-downloads.""" + + async def test_cache_true_reuses_dataset(self): + from datasets import DownloadMode + + loader = ConcreteRemoteLoader() + with patch("pyrit.datasets.seed_datasets.remote.remote_dataset_loader.load_dataset") as mock_load: + await loader._fetch_from_huggingface_async(dataset_name="owner/ds", split="train", cache=True) + + assert mock_load.call_args.kwargs["download_mode"] == DownloadMode.REUSE_DATASET_IF_EXISTS + assert mock_load.call_args.kwargs["cache_dir"] is not None + + async def test_cache_false_forces_redownload(self): + from datasets import DownloadMode + + loader = ConcreteRemoteLoader() + with patch("pyrit.datasets.seed_datasets.remote.remote_dataset_loader.load_dataset") as mock_load: + await loader._fetch_from_huggingface_async(dataset_name="owner/ds", split="train", cache=False) + + assert mock_load.call_args.kwargs["download_mode"] == DownloadMode.FORCE_REDOWNLOAD + assert mock_load.call_args.kwargs["cache_dir"] is None diff --git a/tests/unit/memory/memory_interface/test_interface_seed_prompts.py b/tests/unit/memory/memory_interface/test_interface_seed_prompts.py index 086e35a65b..37d256612d 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_prompts.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_prompts.py @@ -4,10 +4,11 @@ import os import tempfile from collections.abc import Sequence -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from sqlalchemy.exc import SQLAlchemyError from pyrit.memory import MemoryInterface from pyrit.models import MessagePiece, SeedDataset, SeedGroup, SeedObjective, SeedPrompt @@ -1101,3 +1102,116 @@ async def test_get_seed_groups_filter_by_count(sqlite_instance: MemoryInterface) # Test without filtering (should return all) all_groups = sqlite_instance.get_seed_groups() assert len(all_groups) == 2 + + +async def test_replace_seeds_for_dataset_async_replaces_all(sqlite_instance: MemoryInterface): + """replace_seeds_for_dataset_async swaps the target dataset's seeds and leaves others intact.""" + await sqlite_instance.add_seeds_to_memory_async( + seeds=[ + SeedPrompt(value="a1", dataset_name="alpha", data_type="text"), + SeedPrompt(value="a2", dataset_name="alpha", data_type="text"), + SeedPrompt(value="b1", dataset_name="beta", data_type="text"), + ], + added_by="seeding", + ) + + deleted = await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="alpha", + seeds=[SeedPrompt(value="a3", dataset_name="alpha", data_type="text")], + added_by="refresh", + ) + + assert deleted == 2 + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="alpha")} == {"a3"} + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="beta")} == {"b1"} + + +async def test_replace_seeds_for_dataset_async_new_dataset_inserts(sqlite_instance: MemoryInterface): + """Replacing a dataset with no existing rows simply inserts the new seeds and returns 0.""" + deleted = await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="fresh", + seeds=[SeedPrompt(value="v1", dataset_name="fresh", data_type="text")], + added_by="refresh", + ) + + assert deleted == 0 + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="fresh")} == {"v1"} + + +async def test_replace_seeds_for_dataset_async_empty_name_raises(sqlite_instance: MemoryInterface): + """An empty dataset_name is rejected to avoid an accidental mass delete.""" + with pytest.raises(ValueError, match="dataset_name"): + await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="", + seeds=[SeedPrompt(value="v1", dataset_name="x", data_type="text")], + added_by="refresh", + ) + + +async def test_replace_seeds_for_dataset_async_empty_seeds_raises(sqlite_instance: MemoryInterface): + """Refusing empty seeds prevents replacing a dataset with nothing (i.e. wiping it).""" + with pytest.raises(ValueError, match="non-empty"): + await sqlite_instance.replace_seeds_for_dataset_async(dataset_name="alpha", seeds=[], added_by="refresh") + + +async def test_replace_seeds_for_dataset_async_mismatched_name_raises(sqlite_instance: MemoryInterface): + """Seeds tagged for a different dataset are rejected before any delete, avoiding a cross-wipe.""" + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="keep", dataset_name="alpha", data_type="text")], + added_by="seeding", + ) + + with pytest.raises(ValueError, match="mismatched"): + await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="alpha", + seeds=[SeedPrompt(value="foreign", dataset_name="beta", data_type="text")], + added_by="refresh", + ) + + # The guard fires before the delete, so alpha is untouched and beta was never created. + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="alpha")} == {"keep"} + assert sqlite_instance.get_seeds(dataset_name="beta") == [] + + +async def test_replace_seeds_for_dataset_async_mixed_none_and_foreign_name_raises(sqlite_instance: MemoryInterface): + """A mix of a None dataset_name and a foreign one still raises ValueError (not TypeError).""" + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="keep", dataset_name="alpha", data_type="text")], + added_by="seeding", + ) + + seed_unnamed = SeedPrompt(value="unnamed", data_type="text") # dataset_name defaults to None + seed_foreign = SeedPrompt(value="foreign", dataset_name="beta", data_type="text") + + with pytest.raises(ValueError, match="mismatched"): + await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="alpha", + seeds=[seed_unnamed, seed_foreign], + added_by="refresh", + ) + + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="alpha")} == {"keep"} + + +async def test_replace_seeds_for_dataset_async_rolls_back_on_error(sqlite_instance: MemoryInterface): + """A failure during the replace rolls back the delete too, so existing seeds are preserved.""" + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="old", dataset_name="d", data_type="text")], + added_by="seeding", + ) + + real_session = sqlite_instance.get_session() + real_session.commit = MagicMock(side_effect=SQLAlchemyError("commit failed")) + real_session.rollback = MagicMock(side_effect=real_session.rollback) + + with patch.object(sqlite_instance, "get_session", return_value=real_session): + with pytest.raises(SQLAlchemyError, match="commit failed"): + await sqlite_instance.replace_seeds_for_dataset_async( + dataset_name="d", + seeds=[SeedPrompt(value="new", dataset_name="d", data_type="text")], + added_by="refresh", + ) + + real_session.rollback.assert_called_once() + # The delete was rolled back with the failed insert -> the original seed survives. + assert {seed.value for seed in sqlite_instance.get_seeds(dataset_name="d")} == {"old"} diff --git a/tests/unit/setup/test_refresh_datasets.py b/tests/unit/setup/test_refresh_datasets.py new file mode 100644 index 0000000000..0ed3c0568f --- /dev/null +++ b/tests/unit/setup/test_refresh_datasets.py @@ -0,0 +1,417 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Unit tests for the RefreshDatasets initializer. +""" + +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from pyrit.datasets import SeedDatasetProvider +from pyrit.memory import CentralMemory, MemoryInterface +from pyrit.models import SeedDataset, SeedPrompt +from pyrit.setup.initializers.refresh_datasets import RefreshDatasets + + +def _make_dataset(*, dataset_name: str, values: list[str], harm_categories: list[str] | None = None) -> SeedDataset: + seeds = [ + SeedPrompt( + value=value, + dataset_name=dataset_name, + data_type="text", + harm_categories=harm_categories, + ) + for value in values + ] + return SeedDataset(seeds=seeds, name=dataset_name, dataset_name=dataset_name) + + +class TestRefreshDatasetsProperties: + """Property and parameter surface tests.""" + + def test_description_mentions_refresh(self) -> None: + description = RefreshDatasets().description + assert isinstance(description, str) + assert "refresh" in description.lower() + + def test_required_env_vars_is_empty(self) -> None: + assert RefreshDatasets().required_env_vars == [] + + def test_supported_parameters_defaults(self) -> None: + params = {p.name: p for p in RefreshDatasets().supported_parameters} + assert params["days"].default == RefreshDatasets.DEFAULT_DAYS + assert params["dataset_names"].default == [] + assert "tags" not in params + + +class TestRefreshDatasetsParseDays: + """Validation of the days parameter.""" + + def test_default_when_absent(self) -> None: + initializer = RefreshDatasets() + initializer.params = {} + assert initializer._parse_days() == RefreshDatasets.DEFAULT_DAYS + + def test_zero_allowed(self) -> None: + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + assert initializer._parse_days() == 0 + + @pytest.mark.parametrize("bad", [["-1"], ["abc"], ["3.5"], ["1", "2"], []]) + def test_invalid_days(self, bad: list[str]) -> None: + initializer = RefreshDatasets() + # An empty list means "absent" -> default, so only non-empty invalid values raise. + initializer.params = {"days": bad} if bad else {} + if bad: + with pytest.raises(ValueError): + initializer._parse_days() + else: + assert initializer._parse_days() == RefreshDatasets.DEFAULT_DAYS + + +class TestRefreshDatasetsSelection: + """Selection precedence and provider-registration filtering (memory mocked).""" + + def _mock_memory(self, *, names_in_memory: list[str]) -> MagicMock: + memory = MagicMock(spec=MemoryInterface) + memory.get_seed_dataset_names.return_value = names_in_memory + return memory + + async def test_empty_memory_returns_without_fetch(self) -> None: + initializer = RefreshDatasets() + memory = self._mock_memory(names_in_memory=[]) + + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object(SeedDatasetProvider, "fetch_datasets_async", new_callable=AsyncMock) as mock_fetch, + ): + await initializer.initialize_async() + + mock_fetch.assert_not_called() + + async def test_explicit_names_skip_unregistered_and_not_in_memory(self) -> None: + initializer = RefreshDatasets() + initializer.params = {"days": ["0"], "dataset_names": ["in_both", "not_registered", "not_in_memory"]} + memory = self._mock_memory(names_in_memory=["in_both", "not_registered"]) + + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object( + SeedDatasetProvider, + "get_all_dataset_names_async", + new_callable=AsyncMock, + return_value=["in_both", "not_in_memory"], + ), + patch.object(SeedDatasetProvider, "fetch_datasets_async", new_callable=AsyncMock) as mock_fetch, + ): + mock_fetch.return_value = [_make_dataset(dataset_name="in_both", values=["v"])] + + await initializer.initialize_async() + + # Only "in_both" is both in memory and registered. + assert mock_fetch.call_count == 1 + assert mock_fetch.call_args.kwargs["dataset_names"] == ["in_both"] + assert mock_fetch.call_args.kwargs["cache"] is False + + async def test_names_only_consider_registered_and_in_memory(self) -> None: + # dataset_names selection ignores the tag filter entirely (tags are not a parameter). + initializer = RefreshDatasets() + initializer.params = {"days": ["0"], "dataset_names": ["a"]} + memory = self._mock_memory(names_in_memory=["a", "b"]) + + with ( + patch.object(CentralMemory, "get_memory_instance", return_value=memory), + patch.object( + SeedDatasetProvider, + "get_all_dataset_names_async", + new_callable=AsyncMock, + return_value=["a", "b"], + ) as mock_names, + patch.object(SeedDatasetProvider, "fetch_datasets_async", new_callable=AsyncMock) as mock_fetch, + ): + mock_fetch.return_value = [_make_dataset(dataset_name="a", values=["v"])] + await initializer.initialize_async() + + # selection should never consult a tag filter + for call in mock_names.call_args_list: + assert call.kwargs.get("filters") is None + assert mock_fetch.call_args.kwargs["dataset_names"] == ["a"] + + +@pytest.mark.usefixtures("patch_central_database") +class TestRefreshDatasetsStaleness: + """Staleness threshold behavior against a real SQLite memory.""" + + async def _seed(self, memory: MemoryInterface, *, dataset_name: str, days_old: int) -> None: + date_added = datetime.now(tz=timezone.utc) - timedelta(days=days_old) + seed = SeedPrompt(value=f"v-{dataset_name}", dataset_name=dataset_name, data_type="text", date_added=date_added) + await memory.add_seeds_to_memory_async(seeds=[seed], added_by="seeding") + + async def test_days_zero_refreshes_recent_dataset(self, sqlite_instance: MemoryInterface) -> None: + await self._seed(sqlite_instance, dataset_name="fresh", days_old=0) + initializer = RefreshDatasets() + assert initializer._is_stale(memory=sqlite_instance, dataset_name="fresh", days=0) is True + + async def test_recent_dataset_not_stale(self, sqlite_instance: MemoryInterface) -> None: + await self._seed(sqlite_instance, dataset_name="fresh", days_old=1) + initializer = RefreshDatasets() + assert initializer._is_stale(memory=sqlite_instance, dataset_name="fresh", days=30) is False + + async def test_old_dataset_is_stale(self, sqlite_instance: MemoryInterface) -> None: + await self._seed(sqlite_instance, dataset_name="old", days_old=40) + initializer = RefreshDatasets() + assert initializer._is_stale(memory=sqlite_instance, dataset_name="old", days=30) is True + + async def test_cutoff_is_inclusive(self, sqlite_instance: MemoryInterface) -> None: + fixed_now = datetime(2026, 6, 1, 12, 0, 0, tzinfo=timezone.utc) + at_cutoff = fixed_now - timedelta(days=30) + just_newer = at_cutoff + timedelta(microseconds=1) + + seed_at = SeedPrompt(value="at", dataset_name="at_cutoff", data_type="text", date_added=at_cutoff) + seed_new = SeedPrompt(value="new", dataset_name="just_newer", data_type="text", date_added=just_newer) + await sqlite_instance.add_seeds_to_memory_async(seeds=[seed_at, seed_new], added_by="seeding") + + initializer = RefreshDatasets() + with patch("pyrit.setup.initializers.refresh_datasets.datetime") as mock_dt: + mock_dt.now.return_value = fixed_now + assert initializer._is_stale(memory=sqlite_instance, dataset_name="at_cutoff", days=30) is True + assert initializer._is_stale(memory=sqlite_instance, dataset_name="just_newer", days=30) is False + + async def test_no_op_when_all_fresh(self, sqlite_instance: MemoryInterface) -> None: + await self._seed(sqlite_instance, dataset_name="fresh", days_old=1) + initializer = RefreshDatasets() + initializer.params = {"days": ["30"]} + + with ( + patch.object( + SeedDatasetProvider, + "get_all_dataset_names_async", + new_callable=AsyncMock, + return_value=["fresh"], + ), + patch.object(SeedDatasetProvider, "fetch_datasets_async", new_callable=AsyncMock) as mock_fetch, + ): + await initializer.initialize_async() + + mock_fetch.assert_not_called() + + +@pytest.mark.usefixtures("patch_central_database") +class TestRefreshDatasetsRefreshCorrectness: + """End-to-end replace semantics against a real SQLite memory (only the provider is mocked).""" + + async def _run_refresh(self, *, new_dataset: SeedDataset, dataset_name: str) -> None: + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + with ( + patch.object( + SeedDatasetProvider, + "get_all_dataset_names_async", + new_callable=AsyncMock, + return_value=[dataset_name], + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + return_value=[new_dataset], + ), + ): + await initializer.initialize_async() + + async def test_metadata_only_change_replaces_row(self, sqlite_instance: MemoryInterface) -> None: + old = SeedPrompt(value="same-value", dataset_name="d", data_type="text", harm_categories=["oldharm"]) + await sqlite_instance.add_seeds_to_memory_async(seeds=[old], added_by="seeding") + + new_dataset = _make_dataset(dataset_name="d", values=["same-value"], harm_categories=["newharm"]) + await self._run_refresh(new_dataset=new_dataset, dataset_name="d") + + result = sqlite_instance.get_seeds(dataset_name="d") + assert len(result) == 1 + assert result[0].harm_categories == ["newharm"] + + async def test_value_change_replaces_row(self, sqlite_instance: MemoryInterface) -> None: + old = SeedPrompt(value="v1", dataset_name="d", data_type="text") + await sqlite_instance.add_seeds_to_memory_async(seeds=[old], added_by="seeding") + + new_dataset = _make_dataset(dataset_name="d", values=["v2"]) + await self._run_refresh(new_dataset=new_dataset, dataset_name="d") + + result = sqlite_instance.get_seeds(dataset_name="d") + assert len(result) == 1 + assert result[0].value == "v2" + + async def test_upstream_removed_seed_disappears(self, sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[ + SeedPrompt(value="v1", dataset_name="d", data_type="text"), + SeedPrompt(value="v2", dataset_name="d", data_type="text"), + ], + added_by="seeding", + ) + + new_dataset = _make_dataset(dataset_name="d", values=["v1"]) + await self._run_refresh(new_dataset=new_dataset, dataset_name="d") + + values = {seed.value for seed in sqlite_instance.get_seeds(dataset_name="d")} + assert values == {"v1"} + + async def test_other_datasets_untouched(self, sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[ + SeedPrompt(value="keep", dataset_name="other", data_type="text"), + SeedPrompt(value="old", dataset_name="d", data_type="text"), + ], + added_by="seeding", + ) + + new_dataset = _make_dataset(dataset_name="d", values=["new"]) + await self._run_refresh(new_dataset=new_dataset, dataset_name="d") + + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="other")} == {"keep"} + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"new"} + + async def test_failed_fetch_leaves_existing_seeds_intact(self, sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="v1", dataset_name="d", data_type="text")], + added_by="seeding", + ) + + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + with ( + patch.object( + SeedDatasetProvider, "get_all_dataset_names_async", new_callable=AsyncMock, return_value=["d"] + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + side_effect=RuntimeError("network down"), + ), + ): + await initializer.initialize_async() + + # Fetch failed before delete -> the original seed is still present. + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v1"} + + async def test_no_dataset_returned_does_not_wipe_dataset(self, sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="v1", dataset_name="d", data_type="text")], + added_by="seeding", + ) + + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + with ( + patch.object( + SeedDatasetProvider, "get_all_dataset_names_async", new_callable=AsyncMock, return_value=["d"] + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + return_value=[], + ), + ): + await initializer.initialize_async() + + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v1"} + + async def test_empty_dataset_does_not_wipe_dataset(self, sqlite_instance: MemoryInterface) -> None: + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="v1", dataset_name="d", data_type="text")], + added_by="seeding", + ) + + # SeedDataset validation forbids empty seeds, so use a spec'd stand-in to exercise the guard. + empty_dataset = MagicMock(spec=SeedDataset) + empty_dataset.seeds = [] + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + with ( + patch.object( + SeedDatasetProvider, "get_all_dataset_names_async", new_callable=AsyncMock, return_value=["d"] + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + return_value=[empty_dataset], + ), + ): + await initializer.initialize_async() + + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v1"} + + async def test_insert_failure_preserves_existing_seeds(self, sqlite_instance: MemoryInterface) -> None: + # The initializer isolates a failed replace: replace_seeds_for_dataset_async is mocked to + # raise before touching storage, so the existing seeds stay intact and the dataset stays + # selectable for a later retry. (Atomicity of the replace itself -- rolling the delete back + # with a failed insert -- is covered at the memory layer by + # test_replace_seeds_for_dataset_async_rolls_back_on_error.) + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="v1", dataset_name="d", data_type="text")], + added_by="seeding", + ) + new_dataset = _make_dataset(dataset_name="d", values=["v2"]) + + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + + with ( + patch.object( + SeedDatasetProvider, "get_all_dataset_names_async", new_callable=AsyncMock, return_value=["d"] + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + return_value=[new_dataset], + ), + patch.object( + sqlite_instance, + "replace_seeds_for_dataset_async", + new_callable=AsyncMock, + side_effect=RuntimeError("insert failed"), + ), + ): + await initializer.initialize_async() + + # The failed refresh left the original seed untouched (the initializer did not wipe it). + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v1"} + + # The dataset is still in memory, so a later successful run refreshes it. + await self._run_refresh(new_dataset=new_dataset, dataset_name="d") + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v2"} + + async def test_mismatched_dataset_name_does_not_replace(self, sqlite_instance: MemoryInterface) -> None: + # Guard against a provider returning seeds tagged with a different dataset_name, which would + # otherwise delete the requested dataset and insert unrelated seeds under another name. + await sqlite_instance.add_seeds_to_memory_async( + seeds=[SeedPrompt(value="v1", dataset_name="d", data_type="text")], + added_by="seeding", + ) + wrong = _make_dataset(dataset_name="other", values=["x"]) + + initializer = RefreshDatasets() + initializer.params = {"days": ["0"]} + with ( + patch.object( + SeedDatasetProvider, "get_all_dataset_names_async", new_callable=AsyncMock, return_value=["d"] + ), + patch.object( + SeedDatasetProvider, + "fetch_datasets_async", + new_callable=AsyncMock, + return_value=[wrong], + ), + ): + await initializer.initialize_async() + + # The guard rejected the mismatched dataset -> original seeds preserved, nothing leaked. + assert {s.value for s in sqlite_instance.get_seeds(dataset_name="d")} == {"v1"} + assert sqlite_instance.get_seeds(dataset_name="other") == []