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

import logging
from typing import Literal, cast

from pyrit.datasets.seed_datasets.remote.remote_dataset_loader import (
    _RemoteDatasetLoader,
)
from pyrit.models import SeedDataset, SeedPrompt, SeedUnion

logger = logging.getLogger(__name__)


class _DangerousQADataset(_RemoteDatasetLoader):
    """
    Loader for the DangerousQA dataset.

    DangerousQA contains ~200 harmful questions generated by Shaikh et al.
    (2022) from a single seed prompt ("Generate a list of hateful questions").
    The paper describes the dataset as covering racist, stereotypical, sexist,
    illegal, toxic, and harmful content, but these labels apply to the dataset
    as a whole — individual prompts are not categorised in the source release.
    The loader therefore leaves ``harm_categories`` empty on each ``SeedPrompt``
    rather than guess. The dataset is commonly reused as a baseline in
    subsequent red-teaming work (e.g., Red-Eval).

    References:
        - https://github.com/SALT-NLP/chain-of-thought-bias
        - [@shaikh2022second]
    License: MIT

    Warning: This dataset contains harmful questions designed to test LLM safety.
    """

    # Metadata
    # Per-prompt harm categories are intentionally omitted: the source JSON has
    # no per-item labels and the paper only describes the dataset's coverage at
    # the aggregate level, so any class-level harm_categories list would
    # mis-label individual prompts.
    modalities: list[str] = ["text"]
    size: str = "medium"  # ~200 seeds
    tags: set[str] = {"default", "safety"}

    def __init__(
        self,
        *,
        source: str = (
            "https://raw.githubusercontent.com/SALT-NLP/chain-of-thought-bias/"
            "445568d3b73f81a9054f51c739172186d5648157/data/dangerous-q/toxic_outs.json"
        ),
        source_type: Literal["public_url", "file"] = "public_url",
    ) -> None:
        """
        Initialize the DangerousQA dataset loader.

        Args:
            source: URL or path to the DangerousQA JSON file. Defaults to a pinned
                commit of the official SALT-NLP/chain-of-thought-bias repository.
            source_type: The type of source ('public_url' or 'file').
        """
        self.source = source
        self.source_type: Literal["public_url", "file"] = source_type

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

    async def fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
        """
        Fetch DangerousQA dataset and return as SeedDataset.

        Args:
            cache: Whether to cache the fetched dataset. Defaults to True.

        Returns:
            SeedDataset: A SeedDataset containing the DangerousQA questions.

        Raises:
            ValueError: If the source JSON is not a list of strings.
        """
        logger.info(f"Loading DangerousQA dataset from {self.source}")

        # The source JSON is a flat list of strings rather than the list-of-dicts
        # shape most loaders use, but the JSON read/write helpers don't enforce
        # any specific shape, so _fetch_from_url handles fetch and caching uniformly.
        raw = self._fetch_from_url(
            source=self.source,
            source_type=self.source_type,
            cache=cache,
        )

        if not all(isinstance(item, str) for item in raw):
            invalid_types = sorted({type(item).__name__ for item in raw if not isinstance(item, str)})
            raise ValueError(
                f"Expected DangerousQA source to contain a JSON list of strings, got items of types: {invalid_types}"
            )

        questions = cast("list[str]", raw)

        authors = [
            "Omar Shaikh",
            "Hongxin Zhang",
            "William Held",
            "Michael Bernstein",
            "Diyi Yang",
        ]
        groups = [
            "Stanford University",
            "Georgia Institute of Technology",
            "Shanghai Jiao Tong University",
        ]
        description = (
            "DangerousQA contains ~200 harmful questions generated by Shaikh et al. "
            "(2022) in 'On Second Thought, Let's Not Think Step by Step! Bias and "
            "Toxicity in Zero-Shot Reasoning'. The paper describes the set as covering "
            "racist, stereotypical, sexist, illegal, toxic, and harmful content, but "
            "individual prompts are not categorised in the source release. The dataset "
            "is commonly reused as a baseline in subsequent red-teaming work (e.g., "
            "Red-Eval)."
        )

        seed_prompts: list[SeedUnion] = [
            SeedPrompt(
                value=question,
                data_type="text",
                dataset_name=self.dataset_name,
                harm_categories=[],
                description=description,
                source=self.source,
                authors=authors,
                groups=groups,
            )
            for question in questions
        ]

        logger.info(f"Successfully loaded {len(seed_prompts)} prompts from DangerousQA dataset")

        return SeedDataset(seeds=seed_prompts, dataset_name=self.dataset_name)
