Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 29 additions & 1 deletion pyrit/score/float_scale/plagiarism_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,27 @@ def __init__(
metric (PlagiarismMetric): The plagiarism detection metric to use. Defaults to PlagiarismMetric.LCS.
n (int): The n-gram size for n-gram similarity. Defaults to 5.
validator (ScorerPromptValidator | None): Custom validator for the scorer. Defaults to None.

Raises:
ValueError: If ``reference_text`` is not a non-empty string or contains no word tokens,
if ``metric`` is not an instance of PlagiarismMetric, or if ``n`` is not an integer >= 1.
"""
super().__init__(validator=validator or self._DEFAULT_VALIDATOR)
if not isinstance(reference_text, str) or not reference_text.strip():
raise ValueError("reference_text must be a non-empty string.")
if not isinstance(metric, PlagiarismMetric):
raise ValueError(f"metric must be an instance of PlagiarismMetric, got {metric!r}.")
if not isinstance(n, int) or isinstance(n, bool) or n < 1:
raise ValueError(f"n must be an integer >= 1, got {n!r}.")

self.reference_text = reference_text
self.metric = metric
self.n = n

if not self._tokenize(reference_text):
raise ValueError("reference_text must contain at least one word token.")

super().__init__(validator=validator or self._DEFAULT_VALIDATOR)

def _build_identifier(self) -> ComponentIdentifier:
"""
Build the identifier for this scorer.
Expand Down Expand Up @@ -152,6 +166,20 @@ def _plagiarism_score(
metric: PlagiarismMetric = PlagiarismMetric.LCS,
n: int = 5,
) -> float:
"""
Compute word-level similarity after validating the metric and n-gram size.

Returns:
float: The normalized similarity score between 0 and 1.

Raises:
ValueError: If ``metric`` is not a PlagiarismMetric or ``n`` is not an integer >= 1.
"""
if not isinstance(n, int) or isinstance(n, bool) or n < 1:
raise ValueError(f"n must be an integer >= 1, got {n!r}.")
if not isinstance(metric, PlagiarismMetric):
raise ValueError(f"metric must be an instance of PlagiarismMetric, got {metric!r}.")

tokens_response = self._tokenize(response)
tokens_reference = self._tokenize(reference)
response_len = len(tokens_response)
Expand Down
82 changes: 73 additions & 9 deletions tests/unit/score/test_plagiarism_scorer.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

from enum import Enum
from unittest.mock import patch

import pytest
Expand All @@ -11,6 +12,13 @@
from pyrit.score import MessageScorable, PlagiarismMetric, PlagiarismScorer


class _OtherMetric(Enum):
LCS = "lcs"
LEVENSHTEIN = "levenshtein"
JACCARD = "jaccard"
INVALID = "invalid"


@pytest.mark.usefixtures("patch_central_database")
class TestPlagiarismScorer:
"""Test cases for the PlagiarismScorer class."""
Expand All @@ -36,6 +44,68 @@ def test_init_with_custom_parameters(self):
assert scorer.metric == metric
assert scorer.n == n

@pytest.mark.parametrize(
"invalid_reference",
["", " ", "\t\n ", None, 123, [], {}],
)
def test_init_rejects_empty_or_non_string_reference_text(self, invalid_reference):
"""Test initialization rejects empty, whitespace-only, or non-string reference text."""
with pytest.raises(ValueError, match="reference_text must be a non-empty string"):
PlagiarismScorer(reference_text=invalid_reference)

@pytest.mark.parametrize(
"no_token_reference",
["!!!", "???", "---", "... ,,, ;;;", " !@#$%^&*() "],
)
def test_init_rejects_reference_text_without_tokens(self, no_token_reference):
"""Test initialization rejects reference text containing no word tokens."""
with pytest.raises(ValueError, match="reference_text must contain at least one word token"):
PlagiarismScorer(reference_text=no_token_reference)

@pytest.mark.parametrize(
"invalid_n",
[0, -1, -5, 1.5, False, True, "3", None, [3]],
)
def test_init_rejects_invalid_n(self, invalid_n):
"""Test initialization rejects n that is not an integer >= 1 or is a boolean."""
with pytest.raises(ValueError, match=r"n must be an integer >= 1"):
PlagiarismScorer(reference_text="Valid reference text", n=invalid_n)

@pytest.mark.parametrize("valid_n", [1, 2, 5, 10])
def test_init_accepts_valid_boundary_n(self, valid_n):
"""Test initialization accepts positive integer n-gram sizes."""
scorer = PlagiarismScorer(reference_text="Valid reference text", n=valid_n)
assert scorer.n == valid_n

@pytest.mark.parametrize(
"invalid_metric",
["lcs", "levenshtein", "jaccard", "invalid", None, 123, *_OtherMetric],
)
def test_init_rejects_invalid_metric(self, invalid_metric):
"""Test initialization rejects metric that is not an instance of PlagiarismMetric."""
with pytest.raises(ValueError, match="metric must be an instance of PlagiarismMetric"):
PlagiarismScorer(reference_text="Valid reference text", metric=invalid_metric)

@pytest.mark.parametrize("invalid_n", [0, -1, 1.5, False, True, "3", None])
def test_plagiarism_score_rejects_invalid_n(self, invalid_n):
"""Test _plagiarism_score rejects invalid n."""
scorer = PlagiarismScorer(reference_text="Valid reference text")
with pytest.raises(ValueError, match=r"n must be an integer >= 1"):
scorer._plagiarism_score(response="test", reference="test", n=invalid_n)

@pytest.mark.parametrize("invalid_metric", ["lcs", "levenshtein", "jaccard", "invalid", None, 123, *_OtherMetric])
@pytest.mark.parametrize(
("response", "reference"),
[("test", "test"), ("", "test"), ("test", ""), ("different", "test"), ("prefix test suffix", "test")],
)
def test_plagiarism_score_rejects_invalid_metric(
self, *, invalid_metric: object, response: str, reference: str
) -> None:
"""Test _plagiarism_score rejects invalid metric."""
scorer = PlagiarismScorer(reference_text="Valid reference text")
with pytest.raises(ValueError, match="metric must be an instance of PlagiarismMetric"):
scorer._plagiarism_score(response=response, reference=reference, metric=invalid_metric)

async def test_score_async_lcs_metric(self):
"""Test scoring with LCS metric."""
reference_text = "The quick brown fox jumps over the lazy dog"
Expand Down Expand Up @@ -327,15 +397,9 @@ def test_plagiarism_score_empty_texts(self, scorer):
assert score == 0.0

def test_plagiarism_score_invalid_metric(self, scorer):
"""Test plagiarism score with mock invalid metric raises ValueError."""
from unittest.mock import MagicMock

# Create a mock metric that has an invalid value
mock_metric = MagicMock()
mock_metric.value = "invalid"

with pytest.raises(ValueError, match="metric must be 'lcs', 'levenshtein', or 'jaccard'"):
scorer._plagiarism_score("hello", "world", metric=mock_metric)
"""Test plagiarism score rejects an unsupported enum."""
with pytest.raises(ValueError, match="metric must be an instance of PlagiarismMetric"):
scorer._plagiarism_score("hello", "world", metric=_OtherMetric.INVALID)

def test_plagiarism_score_case_insensitive(self, scorer):
"""Test that plagiarism score is case insensitive."""
Expand Down
Loading