From 82304ca0ecc860253e2f84156585f05217f65a7f Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 23 Jul 2026 18:00:39 +0000 Subject: [PATCH] FEAT: add RefreshDatasets initializer to atomically refresh stored datasets Adds an opt-in maintenance initializer that re-fetches datasets already present in memory from their registered providers and atomically replaces their stored seeds, so previously loaded copies pick up upstream changes (standardized harm categories, corrected metadata, live threat-feed updates). - memory: add replace_seeds_for_dataset_async, a single-transaction delete-then-insert so a failed refresh preserves the existing seeds rather than leaving the dataset empty. It rejects seeds whose dataset_name does not match the target to prevent a cross-dataset overwrite. Extract a shared _prepare_seed_for_storage_async helper reused by add_seeds_to_memory_async. - setup: add RefreshDatasets(PyRITInitializer) selecting datasets by staleness (days; 0 refreshes all) and optionally dataset_names, isolating per-dataset failures, and guarding the replace with a fetch-first, name-matching check. - datasets: honor cache=False in the HuggingFace loader by forcing FORCE_REDOWNLOAD so a refresh genuinely bypasses the cache. --- .../remote/remote_dataset_loader.py | 6 +- pyrit/memory/memory_interface.py | 134 +++++- pyrit/setup/initializers/__init__.py | 2 + pyrit/setup/initializers/refresh_datasets.py | 254 +++++++++++ .../datasets/test_remote_dataset_loader.py | 24 + .../test_interface_seed_prompts.py | 116 ++++- tests/unit/setup/test_refresh_datasets.py | 417 ++++++++++++++++++ 7 files changed, 930 insertions(+), 23 deletions(-) create mode 100644 pyrit/setup/initializers/refresh_datasets.py create mode 100644 tests/unit/setup/test_refresh_datasets.py 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") == []