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
15 changes: 6 additions & 9 deletions pyrit/analytics/text_matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,8 @@ def is_match(self, *, target: str, text: str) -> bool:
"""
if not text:
return False
if not target.strip():
return False
if self._ignore_whitespace:
target = target.strip()
text = text.strip()
Expand Down Expand Up @@ -125,27 +127,22 @@ def _calculate_ngram_overlap(self, *, target: str, text: str) -> float:

Returns:
float: A score between 0.0 and 1.0 indicating the proportion of target n-grams
found in the text.
found in the text.
"""
if not text:
return 0.0
if len(target) < self._n:
return 0.0 # Confidence is too low for short targets
return 0.0

target_str = target if self._case_sensitive else target.lower()
text_str = text if self._case_sensitive else text.lower()

# Generate all n-grams from target
target_ngrams = {target_str[i : i + self._n] for i in range(len(target_str) - (self._n - 1))}

# Safety check: if no n-grams were generated, return 0.0
if not target_ngrams:
return 0.0

# Count how many target n-grams are found in text
matching_ngrams = sum(int(ngram in text_str) for ngram in target_ngrams)

# Calculate proportion of matching n-grams
return matching_ngrams / len(target_ngrams)

def get_overlap_score(self, *, target: str, text: str) -> float:
Expand All @@ -155,10 +152,10 @@ def get_overlap_score(self, *, target: str, text: str) -> float:
Useful for getting detailed scoring information.

Args:
target (str): The target string to match.
target (str): The string to search for.
text (str): The text to search in.

Returns:
float: The n-gram overlap score between 0.0 and 1.0.
float: The n-gram overlap score between target and text.
"""
return self._calculate_ngram_overlap(target=target, text=text)
10 changes: 10 additions & 0 deletions tests/unit/analytics/test_text_matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,16 @@ def test_empty_text(self):
matcher = ExactTextMatching(case_sensitive=False)
assert matcher.is_match(target="hello", text="") is False

def test_empty_target(self):
matcher = ExactTextMatching()
Comment thread
romanlutz marked this conversation as resolved.
assert matcher.is_match(target="", text="hello world") is False
assert matcher.is_match(target=" \n ", text="hello world") is False

def test_whitespace_only_target_when_whitespace_is_not_ignored(self):
matcher = ExactTextMatching(ignore_whitespace=False)
assert matcher.is_match(target=" ", text="hello world") is False
assert matcher.is_match(target=" \n\t ", text="hello world") is False

def test_partial_match(self):
matcher = ExactTextMatching(case_sensitive=False)
assert matcher.is_match(target="World", text="Hello World") is True
Expand Down
42 changes: 30 additions & 12 deletions tests/unit/score/test_decoding_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,8 +34,6 @@ def sample_message_pieces():
class TestDecodingScorer:
async def test_decoding_scorer_original_value_match(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces

# Mock memory to return the user request
memory = MagicMock(MemoryInterface)
user_message = Message(message_pieces=[user_piece])
memory.get_request_from_response.return_value = user_message
Expand All @@ -45,7 +43,7 @@ async def test_decoding_scorer_original_value_match(self, patch_central_database
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is True # "secret" is in the response
assert score[0].get_value() is True
assert score[0].score_type == "true_false"
assert score[0].score_category == ["decoding"]

Expand All @@ -62,7 +60,7 @@ async def test_decoding_scorer_converted_value_match(self, patch_central_databas
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is True # "encoded_secret" is in the response
assert score[0].get_value() is True

async def test_decoding_scorer_metadata_match(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
Expand All @@ -77,7 +75,7 @@ async def test_decoding_scorer_metadata_match(self, patch_central_database, samp
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is True # "decoded_secret" from metadata is in response
assert score[0].get_value() is True

async def test_decoding_scorer_no_match(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
Expand All @@ -94,6 +92,30 @@ async def test_decoding_scorer_no_match(self, patch_central_database, sample_mes
assert len(score) == 1
assert score[0].get_value() is False

async def test_decoding_scorer_empty_converted_value_does_not_match(
self, patch_central_database, sample_message_pieces
):
user_piece, _ = sample_message_pieces
user_piece.original_value = "secret"
user_piece.converted_value = ""

assistant_piece = MessagePiece(
role="assistant",
original_value="I refuse",
converted_value="I refuse",
)

memory = MagicMock(MemoryInterface)
user_message = Message(message_pieces=[user_piece])
memory.get_request_from_response.return_value = user_message

with patch.object(CentralMemory, "get_memory_instance", return_value=memory):
scorer = DecodingScorer(categories=["decoding"])
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is False

async def test_decoding_scorer_case_insensitive(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
assistant_piece.converted_value = "Response with SECRET in it"
Expand All @@ -103,12 +125,11 @@ async def test_decoding_scorer_case_insensitive(self, patch_central_database, sa
memory.get_request_from_response.return_value = user_message

with patch.object(CentralMemory, "get_memory_instance", return_value=memory):
# Default is case insensitive
scorer = DecodingScorer(categories=["decoding"])
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is True # Case insensitive match
assert score[0].get_value() is True

async def test_decoding_scorer_case_sensitive(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
Expand All @@ -124,25 +145,22 @@ async def test_decoding_scorer_case_sensitive(self, patch_central_database, samp
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is False # Case sensitive, no match
assert score[0].get_value() is False

async def test_decoding_scorer_approximate_matching(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
# Partial match - has some n-grams in common
assistant_piece.converted_value = "Response with sec ret characters separated"

memory = MagicMock(MemoryInterface)
user_message = Message(message_pieces=[user_piece])
memory.get_request_from_response.return_value = user_message

with patch.object(CentralMemory, "get_memory_instance", return_value=memory):
# Use approximate matching with low threshold
text_matcher = ApproximateTextMatching(threshold=0.3, n=3, case_sensitive=False)
scorer = DecodingScorer(text_matcher=text_matcher, categories=["decoding"])
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
# Should detect partial match with low threshold

async def test_decoding_scorer_approximate_no_match(self, patch_central_database, sample_message_pieces):
user_piece, assistant_piece = sample_message_pieces
Expand All @@ -158,4 +176,4 @@ async def test_decoding_scorer_approximate_no_match(self, patch_central_database
score = await scorer._score_piece_async(assistant_piece)

assert len(score) == 1
assert score[0].get_value() is False # Below threshold
assert score[0].get_value() is False
Loading