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

from collections.abc import Sequence
from typing import get_args

from pyrit.models import ChatMessageRole, Message, MessagePiece, PromptDataType

#: Roles a scorer reads unless it declares otherwise. ``simulated_assistant`` is opt-in
#: because a prepended turn is fabricated history rather than something the target said,
#: and a scorer that judges the target must not mistake one for the other.
DEFAULT_SUPPORTED_ROLES: tuple[ChatMessageRole, ...] = tuple(
    role for role in get_args(ChatMessageRole) if role != "simulated_assistant"
)


class ScorerPromptValidator:
    """
    Validates message pieces and scorer configurations.

    This class provides validation for scorer inputs, ensuring that message pieces meet
    required criteria such as data types, roles, and metadata requirements.
    """

    def __init__(
        self,
        *,
        supported_data_types: Sequence[PromptDataType] | None = None,
        required_metadata: Sequence[str] | None = None,
        supported_roles: Sequence[ChatMessageRole] | None = None,
        max_pieces_in_response: int | None = None,
        max_text_length: int | None = None,
        enforce_all_pieces_valid: bool | None = False,
        raise_on_no_valid_pieces: bool | None = False,
        is_objective_required: bool = False,
    ) -> None:
        """
        Initialize the ScorerPromptValidator.

        Args:
            supported_data_types (Sequence[PromptDataType] | None): Data types that the scorer supports.
                Defaults to all data types if not provided.
            required_metadata (Sequence[str] | None): Metadata keys that must be present in message pieces.
                Defaults to empty list.
            supported_roles (Sequence[ChatMessageRole] | None): Message roles that the scorer reads. Roles are
                compared against the stored role, so ``simulated_assistant`` must be listed to read prepended
                turns. Defaults to every role except ``simulated_assistant``.
            max_pieces_in_response (int | None): Maximum number of pieces allowed in a response.
                Defaults to None (no limit).
            max_text_length (int | None): Maximum character length for text data type pieces.
                Defaults to None (no limit).
            enforce_all_pieces_valid (bool | None): Whether all pieces must be valid or just at least one.
                Defaults to False.
            raise_on_no_valid_pieces (bool | None): Whether to raise ValueError when no pieces are valid.
                Defaults to False, allowing scorers to return ``[]`` for non-applicable evidence.
                Set to True to raise an exception instead.
            is_objective_required (bool): Whether an objective must be provided for scoring. Defaults to False.
        """
        if supported_data_types:
            self._supported_data_types = supported_data_types
        else:
            self._supported_data_types = get_args(PromptDataType)

        if supported_roles:
            self._supported_roles = supported_roles
        else:
            self._supported_roles = DEFAULT_SUPPORTED_ROLES

        self._required_metadata = required_metadata or []

        self._max_pieces_in_response = max_pieces_in_response
        self._max_text_length = max_text_length
        self._enforce_all_pieces_valid = enforce_all_pieces_valid
        self._raise_on_no_valid_pieces = raise_on_no_valid_pieces

        self._is_objective_required = is_objective_required

    @property
    def is_objective_required(self) -> bool:
        """Whether the scorer uses the objective as a required criterion."""
        return self._is_objective_required

    def validate(self, message: Message, objective: str | None) -> None:
        """
        Validate a message and objective against configured requirements.

        Args:
            message (Message): The message to validate.
            objective (str | None): The objective string, if required.

        Raises:
            ValueError: If validation fails due to unsupported pieces, exceeding max pieces, or missing objective.
        """
        valid_pieces_count = 0
        for piece in message.message_pieces:
            if self.is_message_piece_supported(piece):
                valid_pieces_count += 1
            elif self._enforce_all_pieces_valid:
                raise ValueError(
                    f"Message piece {piece.id} with data type {piece.converted_value_data_type} is not supported."
                )

        if valid_pieces_count < 1 and self._raise_on_no_valid_pieces:
            attempted_metadata = [getattr(piece, "prompt_metadata", None) for piece in message.message_pieces]
            raise ValueError(
                "There are no valid pieces to score. \n\n"
                f"Required types: {self._supported_data_types}. "
                f"Required metadata: {self._required_metadata}. "
                f"Length limit: {self._max_pieces_in_response}. "
                f"Objective required: {self._is_objective_required}. "
                f"Message pieces: {message.message_pieces}. "
                f"Prompt metadata: {attempted_metadata}. "
                f"Objective included: {objective}. "
            )

        if self._max_pieces_in_response is not None and len(message.message_pieces) > self._max_pieces_in_response:
            raise ValueError(
                f"Message has {len(message.message_pieces)} pieces, "
                f"exceeding the limit of {self._max_pieces_in_response}."
            )

        if self._is_objective_required and not objective:
            raise ValueError("Objective is required but not provided.")

    def is_role_supported(self, message_piece: MessagePiece) -> bool:
        """
        Check whether this scorer reads pieces in the given piece's role.

        The stored role is compared rather than ``api_role``, so a prepended
        ``simulated_assistant`` turn stays distinguishable from a real response.

        Args:
            message_piece (MessagePiece): The message piece to check.

        Returns:
            bool: True if the scorer reads this role.
        """
        return message_piece.role in self._supported_roles

    def is_message_piece_supported(self, message_piece: MessagePiece) -> bool:
        """
        Check if a message piece is supported by this validator.

        Args:
            message_piece (MessagePiece): The message piece to check.

        Returns:
            bool: True if the message piece meets all validation criteria, False otherwise.
        """
        if message_piece.converted_value_data_type not in self._supported_data_types:
            return False

        for metadata in self._required_metadata:
            if metadata not in message_piece.prompt_metadata:
                return False

        if not self.is_role_supported(message_piece):
            return False

        # Check text length limit for text data types
        if self._max_text_length is not None and message_piece.converted_value_data_type == "text":
            text_length = len(message_piece.converted_value) if message_piece.converted_value else 0
            if text_length > self._max_text_length:
                return False

        return True
