"""Tests for dataset loading and synthetic generation.""" from __future__ import annotations import pytest from data_cleaning_env.datasets import ( TASK_CONFIGS, generate_synthetic_dataset, load_clean_dataset, ) class TestDatasetLoading: def test_four_task_configs(self) -> None: assert len(TASK_CONFIGS) == 4 assert "expert" in TASK_CONFIGS @pytest.mark.parametrize("task", ["easy", "medium", "hard", "expert"]) def test_load_returns_nonempty(self, task: str) -> None: df, target = load_clean_dataset(task) assert len(df) > 0 assert target in df.columns def test_unknown_task_raises(self) -> None: with pytest.raises(ValueError, match="Unknown task"): load_clean_dataset("nonexistent") class TestSyntheticDataset: def test_default_shape(self) -> None: df, target = generate_synthetic_dataset() assert df.shape[0] == 5000 assert target == "class" assert "class" in df.columns def test_deterministic(self) -> None: df1, _ = generate_synthetic_dataset(seed=99) df2, _ = generate_synthetic_dataset(seed=99) assert df1.equals(df2)