From bcc972a950b71c3a5678a323d11fd1f2bec0a92a Mon Sep 17 00:00:00 2001 From: biefan <70761325+biefan@users.noreply.github.com> Date: Sun, 16 Aug 2026 07:04:46 +0000 Subject: [PATCH] FIX: Offload local dataset file reads --- .../local/local_dataset_loader.py | 15 ++++- .../datasets/test_local_dataset_loader.py | 56 ++++++++++++++++++- 2 files changed, 67 insertions(+), 4 deletions(-) diff --git a/pyrit/datasets/seed_datasets/local/local_dataset_loader.py b/pyrit/datasets/seed_datasets/local/local_dataset_loader.py index af2d427504..92f85ccd12 100644 --- a/pyrit/datasets/seed_datasets/local/local_dataset_loader.py +++ b/pyrit/datasets/seed_datasets/local/local_dataset_loader.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import asyncio import logging from collections.abc import Callable from dataclasses import fields @@ -68,7 +69,7 @@ async def fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset: """ try: logger.info(f"Loading local dataset from {self.file_path}") - dataset = SeedDataset.from_yaml_file(self.file_path) + dataset = await asyncio.to_thread(SeedDataset.from_yaml_file, self.file_path) if not dataset.dataset_name: dataset.dataset_name = self.dataset_name return dataset @@ -91,8 +92,7 @@ async def _parse_metadata_async(self) -> SeedDatasetMetadata | None: """ valid_fields = [f.name for f in fields(SeedDatasetMetadata)] try: - with open(self.file_path, encoding="utf-8") as f: - dataset = yaml.safe_load(f) + dataset = await asyncio.to_thread(self._read_yaml) except Exception as e: logger.error(f"Failed to load local dataset from {self.file_path}: {e}") raise @@ -111,6 +111,15 @@ async def _parse_metadata_async(self) -> SeedDatasetMetadata | None: SeedDatasetMetadata._validate_singular_fields(metadata=result) return result + def _read_yaml(self) -> Any: + """ + Read and parse the local dataset YAML file. + + Returns: + Any: Parsed YAML content. + """ + return yaml.safe_load(self.file_path.read_text(encoding="utf-8")) + def _register_local_datasets() -> None: """ diff --git a/tests/unit/datasets/test_local_dataset_loader.py b/tests/unit/datasets/test_local_dataset_loader.py index 6c51b6b877..66cf663911 100644 --- a/tests/unit/datasets/test_local_dataset_loader.py +++ b/tests/unit/datasets/test_local_dataset_loader.py @@ -2,11 +2,12 @@ # Licensed under the MIT license. from pathlib import Path +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pyrit.datasets.seed_datasets.local.local_dataset_loader import _LocalDatasetLoader -from pyrit.models import SeedDataset +from pyrit.models import SeedDataset, SeedPrompt class TestLocalDatasetLoader: @@ -49,6 +50,59 @@ async def test_fetch_dataset(self, tmp_path, valid_yaml_content): assert len(dataset.prompts) == 1 assert dataset.prompts[0].value == "test prompt" + async def test_fetch_dataset_offloads_file_read(self, tmp_path: Path) -> None: + """Dataset file loading runs outside the event loop thread.""" + file_path = tmp_path / "test.yaml" + loader = _LocalDatasetLoader.__new__(_LocalDatasetLoader) + loader.file_path = file_path + loader._dataset_name = "test_dataset" + expected = SeedDataset( + dataset_name="test_dataset", + seeds=[SeedPrompt(value="test prompt", data_type="text")], + ) + to_thread_mock = AsyncMock(return_value=expected) + + with ( + patch.object(SeedDataset, "from_yaml_file") as load_mock, + patch( + "pyrit.datasets.seed_datasets.local.local_dataset_loader.asyncio.to_thread", + new=to_thread_mock, + ), + ): + dataset = await loader.fetch_dataset_async() + + assert dataset is expected + to_thread_mock.assert_awaited_once_with(load_mock, file_path) + load_mock.assert_not_called() + + async def test_parse_metadata_offloads_file_read(self, tmp_path: Path) -> None: + """Metadata YAML parsing runs outside the event loop thread.""" + file_path = tmp_path / "test.yaml" + loader = _LocalDatasetLoader.__new__(_LocalDatasetLoader) + loader.file_path = file_path + loader._dataset_name = "test_dataset" + read_yaml_mock = MagicMock( + return_value={ + "dataset_name": "test_dataset", + "harm_categories": ["violence"], + } + ) + to_thread_mock = AsyncMock(return_value=read_yaml_mock.return_value) + + with ( + patch.object(loader, "_read_yaml", new=read_yaml_mock), + patch( + "pyrit.datasets.seed_datasets.local.local_dataset_loader.asyncio.to_thread", + new=to_thread_mock, + ), + ): + metadata = await loader._parse_metadata_async() + + assert metadata is not None + assert metadata.harm_categories == {"violence"} + to_thread_mock.assert_awaited_once_with(read_yaml_mock) + read_yaml_mock.assert_not_called() + async def test_fetch_dataset_file_not_found(self): loader = _LocalDatasetLoader(file_path=Path("non_existent.yaml")) with pytest.raises(Exception):