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


from pyrit.analytics.text_matching import ExactTextMatching, TextMatching
from pyrit.models import ComponentIdentifier, MessagePiece, Score
from pyrit.score.scorer_prompt_validator import ScorerPromptValidator
from pyrit.score.true_false.true_false_score_aggregator import (
    TrueFalseAggregatorFunc,
    TrueFalseScoreAggregator,
)
from pyrit.score.true_false.true_false_scorer import MessageTrueFalseScorer


class SubStringScorer(MessageTrueFalseScorer):
    """
    Scorer that checks if a given substring is present in the text.

    This scorer performs substring matching using a configurable text matching strategy.
    Supports both exact substring matching and approximate matching.
    """

    _DEFAULT_VALIDATOR: ScorerPromptValidator = ScorerPromptValidator(supported_data_types=["text"])

    def __init__(
        self,
        *,
        substring: str,
        text_matcher: TextMatching | None = None,
        categories: list[str] | None = None,
        aggregator: TrueFalseAggregatorFunc = TrueFalseScoreAggregator.OR,
        validator: ScorerPromptValidator | None = None,
    ) -> None:
        """
        Initialize the SubStringScorer.

        Args:
            substring (str): The substring to search for in the text.
            text_matcher (TextMatching | None): The text matching strategy to use.
                Defaults to ExactTextMatching with case_sensitive=False.
            categories (list[str] | None): Optional list of categories for the score. Defaults to None.
            aggregator (TrueFalseAggregatorFunc): The aggregator function to use.
                Defaults to TrueFalseScoreAggregator.OR.
            validator (ScorerPromptValidator | None): Custom validator. Defaults to None.
        """
        self._substring = substring
        self._text_matcher = text_matcher if text_matcher else ExactTextMatching(case_sensitive=False)
        self._score_categories = categories if categories else []

        super().__init__(score_aggregator=aggregator, validator=validator or self._DEFAULT_VALIDATOR)

    def _build_identifier(self) -> ComponentIdentifier:
        """
        Build the identifier for this scorer.

        Returns:
            ComponentIdentifier: The identifier for this scorer.
        """
        return self._create_identifier(
            params={
                "substring": self._substring,
                "text_matcher": self._text_matcher.__class__.__name__,
            },
            score_aggregator=self._score_aggregator.__name__,  # type: ignore[ty:unresolved-attribute]
        )

    async def _score_piece_async(self, message_piece: MessagePiece, *, objective: str | None = None) -> list[Score]:
        """
        Score the given message piece based on presence of the substring.

        Args:
            message_piece (MessagePiece): The message piece to score.
            objective (str | None): The objective to evaluate against. Defaults to None.
                Currently not used for this scorer.

        Returns:
            list[Score]: A list containing a single Score object with a boolean value indicating
                whether the substring matches the text according to the matching strategy.
        """
        substring_present = self._text_matcher.is_match(target=self._substring, text=message_piece.converted_value)

        return [
            Score(
                score_value=str(substring_present),
                score_value_description="",
                score_metadata=None,
                score_type="true_false",
                score_category=self._score_categories,
                score_rationale="",
                scorer_class_identifier=self.get_identifier(),
                message_piece_id=message_piece.id,
                objective=objective,
            )
        ]
