Skip to content

Commit 71eaa26

Browse files
rlundeen2Copilot
andauthored
TEST: Moving dataset tests to end-to-end (#1589)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent bd484a0 commit 71eaa26

3 files changed

Lines changed: 184 additions & 33 deletions

File tree

Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,114 @@
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")
Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT license.
3+
4+
"""
5+
Integration test for the LoadDefaultDatasets initializer.
6+
7+
Runs the full pipeline: discovers scenario default datasets, fetches them
8+
from real remote sources, and stores them in in-memory CentralMemory.
9+
"""
10+
11+
import logging
12+
13+
import pytest
14+
15+
from pyrit.memory import CentralMemory
16+
from pyrit.setup.initializers.scenarios.load_default_datasets import LoadDefaultDatasets
17+
18+
logger = logging.getLogger(__name__)
19+
20+
21+
class TestLoadDefaultDatasetsIntegration:
22+
"""Integration test that LoadDefaultDatasets loads real datasets into memory."""
23+
24+
@pytest.mark.asyncio
25+
async def test_initialize_loads_datasets_into_memory(self):
26+
"""
27+
Verify that LoadDefaultDatasets.initialize_async() successfully fetches
28+
real datasets and stores them in CentralMemory.
29+
"""
30+
initializer = LoadDefaultDatasets()
31+
await initializer.initialize_async()
32+
33+
memory = CentralMemory.get_memory_instance()
34+
dataset_names = memory.get_seed_dataset_names()
35+
36+
assert len(dataset_names) > 0, "No datasets were loaded into memory"
37+
logger.info(f"LoadDefaultDatasets loaded {len(dataset_names)} datasets into memory")

tests/integration/datasets/test_seed_dataset_provider_integration.py

Lines changed: 33 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
from pyrit.datasets import SeedDatasetProvider
1212
from pyrit.datasets.seed_datasets.local.local_dataset_loader import _LocalDatasetLoader
13-
from pyrit.datasets.seed_datasets.remote import _VLSUMultimodalDataset
13+
from pyrit.datasets.seed_datasets.remote import _SimpleSafetyTestsDataset, _XSTestDataset
1414
from pyrit.datasets.seed_datasets.seed_metadata import (
1515
SeedDatasetFilter,
1616
)
@@ -19,49 +19,49 @@
1919
logger = logging.getLogger(__name__)
2020

2121

22-
def get_dataset_providers():
23-
"""Helper to get all registered providers for parameterization."""
24-
providers = SeedDatasetProvider.get_all_providers()
25-
return [(name, cls) for name, cls in providers.items()]
22+
# Smoke-test providers covering the three distinct fetch paths:
23+
# - local YAML (no network)
24+
# - remote URL-based (_fetch_from_url via GitHub)
25+
# - remote HuggingFace (_fetch_from_huggingface)
26+
_all_providers = SeedDatasetProvider.get_all_providers()
27+
_SMOKE_PROVIDERS: list[tuple[str, type]] = [
28+
("LocalDataset_access_shell_commands", _all_providers["LocalDataset_access_shell_commands"]),
29+
("_XSTestDataset", _XSTestDataset),
30+
("_SimpleSafetyTestsDataset", _SimpleSafetyTestsDataset),
31+
]
2632

2733

28-
class TestSeedDatasetProviderIntegration:
29-
"""Integration tests for SeedDatasetProvider."""
34+
class TestSeedDatasetSmoke:
35+
"""Smoke tests for a small representative set of dataset providers.
36+
37+
The exhaustive test over all providers lives in tests/end_to_end/test_all_datasets.py.
38+
"""
3039

3140
@pytest.mark.asyncio
32-
@pytest.mark.parametrize("name,provider_cls", get_dataset_providers())
33-
async def test_fetch_dataset_integration(self, name, provider_cls):
41+
@pytest.mark.parametrize("name,provider_cls", _SMOKE_PROVIDERS, ids=[p[0] for p in _SMOKE_PROVIDERS])
42+
async def test_fetch_dataset_smoke(self, name, provider_cls):
3443
"""
35-
Integration test to verify that a specific registered dataset can be fetched.
44+
Verify that a representative provider can be fetched successfully.
3645
37-
This test is parameterized to run for each registered provider.
38-
It verifies that:
39-
1. The dataset can be downloaded/loaded without error
40-
2. The result is a SeedDataset
41-
3. The dataset is not empty (has seeds)
46+
Covers one local, one URL-remote, and one HuggingFace-remote provider
47+
to catch regressions in each fetch path without downloading all 58 datasets.
4248
"""
43-
logger.info(f"Testing provider: {name}")
49+
logger.info(f"Smoke testing provider: {name}")
4450

45-
try:
46-
# Use max_examples for slow providers that fetch many remote images
47-
provider = provider_cls(max_examples=6) if provider_cls == _VLSUMultimodalDataset else provider_cls()
48-
dataset = await provider.fetch_dataset(cache=False)
51+
provider = provider_cls()
52+
dataset = await provider.fetch_dataset(cache=False)
4953

50-
assert isinstance(dataset, SeedDataset), f"{name} did not return a SeedDataset"
51-
assert len(dataset.seeds) > 0, f"{name} returned an empty dataset"
52-
assert dataset.dataset_name, f"{name} has no dataset_name"
54+
assert isinstance(dataset, SeedDataset), f"{name} did not return a SeedDataset"
55+
assert len(dataset.seeds) > 0, f"{name} returned an empty dataset"
56+
assert dataset.dataset_name, f"{name} has no dataset_name"
5357

54-
# Verify seeds have required fields
55-
for seed in dataset.seeds:
56-
assert seed.value, f"Seed in {name} has no value"
57-
assert seed.dataset_name == dataset.dataset_name, (
58-
f"Seed dataset_name mismatch in {name}: {seed.dataset_name} != {dataset.dataset_name}"
59-
)
60-
61-
logger.info(f"Successfully verified {name} with {len(dataset.seeds)} seeds")
58+
for seed in dataset.seeds:
59+
assert seed.value, f"Seed in {name} has no value"
60+
assert seed.dataset_name == dataset.dataset_name, (
61+
f"Seed dataset_name mismatch in {name}: {seed.dataset_name} != {dataset.dataset_name}"
62+
)
6263

63-
except Exception as e:
64-
pytest.fail(f"Failed to fetch dataset from {name}: {str(e)}")
64+
logger.info(f"Smoke test passed for {name} with {len(dataset.seeds)} seeds")
6565

6666

6767
class TestRemoteFilteringIntegration:

0 commit comments

Comments
 (0)