# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import asyncio
import logging
from collections.abc import Callable
from dataclasses import fields
from pathlib import Path
from typing import Any

import yaml

from pyrit.datasets.seed_datasets.seed_dataset_provider import SeedDatasetProvider
from pyrit.datasets.seed_datasets.seed_metadata import (
    SeedDatasetMetadata,
)
from pyrit.models import SeedDataset

logger = logging.getLogger(__name__)


class _LocalDatasetLoader(SeedDatasetProvider):
    """
    Loader for local YAML dataset files.

    This loader discovers and loads datasets from local YAML files.
    Each YAML file should be in the standard SeedDataset format.
    """

    should_register = False

    def __init__(self, *, file_path: Path) -> None:
        """
        Initialize the local dataset loader.

        Args:
            file_path: Path to the YAML dataset file.
        """
        self.file_path = file_path

        # Pre-load to get dataset name
        try:
            dataset = SeedDataset.from_yaml_file(file_path)
            # Use the dataset_name from the YAML if available, otherwise use filename
            self._dataset_name: str = (
                getattr(dataset, "dataset_name", None) or getattr(dataset, "name", None) or file_path.stem
            )
        except Exception as e:
            logger.warning(f"Could not pre-load dataset from {file_path}: {e}")
            self._dataset_name = file_path.stem

    @property
    def dataset_name(self) -> str:
        """The dataset name."""
        return self._dataset_name

    async def fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
        """
        Load the dataset from the local YAML file.

        Args:
            cache: Ignored for local datasets (included for interface consistency).

        Returns:
            SeedDataset: The loaded dataset.

        Raises:
            Exception: If the dataset cannot be loaded.
        """
        try:
            logger.info(f"Loading local dataset from {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
        except Exception as e:
            logger.error(f"Failed to load local dataset from {self.file_path}: {e}")
            raise

    async def _parse_metadata_async(self) -> SeedDatasetMetadata | None:
        """
        Extract metadata from a local YAML file and coerce raw values into typed schema fields.

        YAML produces raw Python primitives (str, list) that must be converted to the
        enum and set types expected by SeedDatasetMetadata before _match_filter can work.

        Returns:
            SeedDatasetMetadata | None: Parsed metadata if available, otherwise None.

        Raises:
            Exception: If the dataset file cannot be read.
        """
        valid_fields = [f.name for f in fields(SeedDatasetMetadata)]
        try:
            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

        if not isinstance(dataset, dict):
            return None

        raw = {k: v for k, v in dataset.items() if k in valid_fields}
        if not raw:
            return None

        coerced = SeedDatasetMetadata._coerce_metadata_values(raw_metadata=raw)
        result = SeedDatasetMetadata(**coerced)
        # Validation after coercion: raw values are strings/lists, not sets.
        # _validate_singular_fields needs sets to check cardinality.
        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:
    """
    Auto-discover and register all YAML files from the seed_datasets directory.
    """
    # Get the path to the seed_datasets directory (parent of this file)
    seed_datasets_path = Path(__file__).parent

    if seed_datasets_path.exists():
        # Search for both .prompt and .yaml files
        for pattern in ["**/*.prompt", "**/*.yaml"]:
            for yaml_file in seed_datasets_path.glob(pattern):
                try:
                    # Create a dynamic subclass for each file to register it
                    # The class name needs to be unique
                    class_name = f"LocalDataset_{yaml_file.stem.replace('-', '_').replace(' ', '_')}"

                    # Define the class dynamically
                    # We set should_register=True so it gets registered
                    # We override __init__ to pass the specific file_path

                    def make_init(path: Path) -> Callable[[Any], None]:
                        def __init__(self: Any) -> None:  # noqa: N807
                            super(self.__class__, self).__init__(file_path=path)

                        return __init__

                    type(
                        class_name,
                        (_LocalDatasetLoader,),
                        {"__init__": make_init(yaml_file), "should_register": True, "__module__": __name__},
                    )

                    logger.debug(f"Registered local dataset loader: {class_name} for {yaml_file.name}")
                except Exception as e:
                    logger.warning(f"Failed to register local dataset {yaml_file}: {e}")
    else:
        logger.warning(f"Seed datasets directory not found: {seed_datasets_path}")


# Execute registration
_register_local_datasets()
