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

import abc
from typing import Any, Generic, Literal, TypeVar

from pydantic import BaseModel

from pyrit.models import ComponentIdentifier, Identifiable, Message

# Type alias for system message handling strategies
SystemMessageBehavior = Literal["keep", "squash", "ignore"]
"""
How to handle system messages in models with varying support:
- "keep": Keep system messages as-is (default for most models)
- "squash": Merge system messages into the following user message
- "ignore": Drop system messages entirely
"""


T = TypeVar("T", bound=BaseModel)


class MessageListNormalizer(abc.ABC, Generic[T]):
    """
    Abstract base class for normalizers that return a list of items.

    Subclasses specify the type T (e.g., Message, ChatMessage) that the list contains.
    """

    @abc.abstractmethod
    async def normalize_async(self, messages: list[Message]) -> list[T]:
        """
        Normalize the list of messages into a list of items.

        Args:
            messages: The list of Message objects to normalize.

        Returns:
            A list of normalized items of type T.

        Note:
            Output metadata is authoritative. A normalizer that creates replacement
            ``Message`` or ``MessagePiece`` objects must preserve any
            ``prompt_metadata`` required by downstream consumers. The target stamps
            the active conversation ID but does not merge removed metadata back in.
        """

    async def normalize_to_dicts_async(self, messages: list[Message]) -> list[dict[str, Any]]:
        """
        Normalize the list of messages into a list of dictionaries.

        This method uses normalize_async and serializes each item with
        ``model_dump(exclude_none=True)``.

        Args:
            messages: The list of Message objects to normalize.

        Returns:
            A list of dictionaries representing the normalized messages.
        """
        normalized = await self.normalize_async(messages)
        return [item.model_dump(exclude_none=True) for item in normalized]


class MessageStringNormalizer(Identifiable, abc.ABC):
    """
    Abstract base class for normalizers that return a string representation.

    Use this for formatting messages into text for non-chat targets or context strings.
    """

    @abc.abstractmethod
    async def normalize_string_async(self, messages: list[Message]) -> str:
        """
        Normalize the list of messages into a string representation.

        Args:
            messages: The list of Message objects to normalize.

        Returns:
            A string representation of the messages.
        """

    def _build_identifier(self) -> ComponentIdentifier:
        """
        Build the formatter's behavioral identifier.

        Stateless custom formatters receive class identity by default. Formatters
        with behavior-changing configuration should override this method.

        Returns:
            The formatter's behavioral identifier.
        """
        return ComponentIdentifier.of(self)


async def apply_system_message_behavior_async(
    messages: list[Message], behavior: SystemMessageBehavior
) -> list[Message]:
    """
    Apply a system message behavior to a list of messages.

    This is a helper function used by normalizers to preprocess messages
    based on how the target handles system messages.

    Args:
        messages: The list of Message objects to process.
        behavior: How to handle system messages:
            - "keep": Return messages unchanged
            - "squash": Merge system messages into the following user message
            - "ignore": Remove system messages

    Returns:
        The processed list of Message objects.

    Raises:
        ValueError: If an unknown behavior is provided.
    """
    if behavior == "keep":
        return messages
    if behavior == "squash":
        # Import here to avoid circular imports
        from pyrit.message_normalizer.generic_system_squash import (
            GenericSystemSquashNormalizer,
        )

        return await GenericSystemSquashNormalizer().normalize_async(messages)
    if behavior == "ignore":
        return [msg for msg in messages if msg.api_role != "system"]
    # This should never happen due to Literal type, but handle it gracefully
    raise ValueError(f"Unknown system message behavior: {behavior}")
