|
| 1 | +# Copyright (c) Microsoft Corporation. |
| 2 | +# Licensed under the MIT license. |
| 3 | + |
| 4 | +""" |
| 5 | +End-to-end tests that verify every registered dataset provider can be fetched. |
| 6 | +
|
| 7 | +These tests download real data from HuggingFace and GitHub, are slow, and are |
| 8 | +subject to transient network failures. They are intended to run daily in e2e CI, |
| 9 | +not on every PR. |
| 10 | +
|
| 11 | +Resiliency: each fetch is retried up to 3 times with exponential backoff to |
| 12 | +handle transient HuggingFace / GitHub rate-limiting and network errors. |
| 13 | +""" |
| 14 | + |
| 15 | +import asyncio |
| 16 | +import logging |
| 17 | +import os |
| 18 | + |
| 19 | +import pytest |
| 20 | +from tenacity import retry, retry_if_exception_type, stop_after_attempt, wait_exponential |
| 21 | + |
| 22 | +from pyrit.datasets import SeedDatasetProvider |
| 23 | +from pyrit.datasets.seed_datasets.remote import ( |
| 24 | + _HarmBenchMultimodalDataset, |
| 25 | + _PromptIntelDataset, |
| 26 | + _VLSUMultimodalDataset, |
| 27 | +) |
| 28 | +from pyrit.models import SeedDataset |
| 29 | +from pyrit.setup import IN_MEMORY, initialize_pyrit_async |
| 30 | + |
| 31 | +logger = logging.getLogger(__name__) |
| 32 | + |
| 33 | +# Per-test timeout in seconds (5 minutes per dataset) |
| 34 | +_TEST_TIMEOUT = 300 |
| 35 | + |
| 36 | +# Transient error types that warrant a retry |
| 37 | +_RETRYABLE_ERRORS = (OSError, ConnectionError, TimeoutError) |
| 38 | + |
| 39 | +# Providers that download many remote images; each image fetch may fail |
| 40 | +# due to rate-limiting, so an empty result is expected in some environments. |
| 41 | +_IMAGE_FETCHING_PROVIDERS: set[type] = {_HarmBenchMultimodalDataset, _VLSUMultimodalDataset} |
| 42 | + |
| 43 | + |
| 44 | +def get_dataset_providers(): |
| 45 | + """Helper to get all registered providers for parameterization.""" |
| 46 | + providers = SeedDatasetProvider.get_all_providers() |
| 47 | + return [(name, cls) for name, cls in providers.items()] |
| 48 | + |
| 49 | + |
| 50 | +@retry( |
| 51 | + stop=stop_after_attempt(3), |
| 52 | + wait=wait_exponential(multiplier=5, min=5, max=60), |
| 53 | + retry=retry_if_exception_type(_RETRYABLE_ERRORS), |
| 54 | + reraise=True, |
| 55 | +) |
| 56 | +async def _fetch_with_retry(provider) -> SeedDataset: |
| 57 | + """Fetch a dataset with retry on transient network errors.""" |
| 58 | + return await provider.fetch_dataset(cache=False) |
| 59 | + |
| 60 | + |
| 61 | +@pytest.fixture(scope="module", autouse=True) |
| 62 | +def _init_memory(): |
| 63 | + """Multimodal providers need CentralMemory to save downloaded images.""" |
| 64 | + asyncio.run(initialize_pyrit_async(memory_db_type=IN_MEMORY)) |
| 65 | + |
| 66 | + |
| 67 | +class TestAllDatasets: |
| 68 | + """Exhaustive test that every registered dataset provider can be fetched.""" |
| 69 | + |
| 70 | + @pytest.mark.asyncio |
| 71 | + @pytest.mark.timeout(_TEST_TIMEOUT) |
| 72 | + @pytest.mark.parametrize("name,provider_cls", get_dataset_providers()) |
| 73 | + async def test_fetch_dataset(self, name, provider_cls): |
| 74 | + """ |
| 75 | + Verify that a specific registered dataset can be fetched. |
| 76 | +
|
| 77 | + This test is parameterized to run for each registered provider. |
| 78 | + It verifies that: |
| 79 | + 1. The dataset can be downloaded/loaded without error |
| 80 | + 2. The result is a SeedDataset |
| 81 | + 3. The dataset is not empty (has seeds) |
| 82 | +
|
| 83 | + Retries up to 3 times on transient network errors. |
| 84 | + """ |
| 85 | + # Skip providers that require credentials not available in CI |
| 86 | + if provider_cls == _PromptIntelDataset and not os.environ.get("PROMPTINTEL_API_KEY"): |
| 87 | + pytest.skip("PROMPTINTEL_API_KEY not set") |
| 88 | + |
| 89 | + logger.info(f"Testing provider: {name}") |
| 90 | + |
| 91 | + try: |
| 92 | + # Limit examples for slow multimodal providers that fetch many remote images |
| 93 | + provider = provider_cls(max_examples=6) if provider_cls == _VLSUMultimodalDataset else provider_cls() |
| 94 | + |
| 95 | + dataset = await _fetch_with_retry(provider) |
| 96 | + except Exception as e: |
| 97 | + # Multimodal providers silently skip failed image downloads. When ALL |
| 98 | + # images fail the resulting empty seed list triggers "SeedDataset cannot |
| 99 | + # be empty". That is a transient environment issue, not a code bug. |
| 100 | + if provider_cls in _IMAGE_FETCHING_PROVIDERS and "cannot be empty" in str(e): |
| 101 | + pytest.skip(f"{name}: all image downloads failed ({e})") |
| 102 | + pytest.fail(f"Failed to fetch dataset from {name}: {e}") |
| 103 | + |
| 104 | + assert isinstance(dataset, SeedDataset), f"{name} did not return a SeedDataset" |
| 105 | + assert dataset.dataset_name, f"{name} has no dataset_name" |
| 106 | + assert len(dataset.seeds) > 0, f"{name} returned an empty dataset" |
| 107 | + |
| 108 | + for seed in dataset.seeds: |
| 109 | + assert seed.value, f"Seed in {name} has no value" |
| 110 | + assert seed.dataset_name == dataset.dataset_name, ( |
| 111 | + f"Seed dataset_name mismatch in {name}: {seed.dataset_name} != {dataset.dataset_name}" |
| 112 | + ) |
| 113 | + |
| 114 | + logger.info(f"Successfully verified {name} with {len(dataset.seeds)} seeds") |
0 commit comments