diff --git a/pyrit/score/true_false/prompt_shield_scorer.py b/pyrit/score/true_false/prompt_shield_scorer.py index 3f0a8d1cc5..5022d0963e 100644 --- a/pyrit/score/true_false/prompt_shield_scorer.py +++ b/pyrit/score/true_false/prompt_shield_scorer.py @@ -83,18 +83,18 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st # Whether or not any of the documents or userPrompt got flagged as an attack result: bool = any(self._parse_response_to_boolean_list(response)) - # Store parsed response in metadata for structured access - try: - meta = json.loads(response) - except Exception: - meta = {"raw": response} + # Store the endpoint response in metadata for structured access. Score only + # accepts flat str/int/float metadata values, so the body is kept as the JSON + # text it arrived as; json.loads on it, as _parse_response_to_boolean_list's + # docstring says. + meta = {"response": response} score = Score( score_type="true_false", score_value=str(result), score_value_description="True if an attack or jailbreak has been detected, else False.", score_category=["attack_detection"], - score_metadata=meta, # type: ignore[ty:invalid-argument-type] + score_metadata=meta, score_rationale="", scorer_class_identifier=self.get_identifier(), message_piece_id=message_piece.id, diff --git a/tests/unit/score/test_prompt_shield_scorer.py b/tests/unit/score/test_prompt_shield_scorer.py index 846147bb5d..e84eeaa385 100644 --- a/tests/unit/score/test_prompt_shield_scorer.py +++ b/tests/unit/score/test_prompt_shield_scorer.py @@ -1,13 +1,16 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import json +import uuid from collections.abc import MutableSequence -from unittest.mock import Mock +from unittest.mock import AsyncMock, MagicMock, Mock import pytest -from unit.mocks import get_sample_conversations +from unit.mocks import get_mock_target_identifier, get_sample_conversations -from pyrit.models import MessagePiece, flatten_to_message_pieces +from pyrit.models import Message, MessagePiece, flatten_to_message_pieces +from pyrit.prompt_target import PromptTarget from pyrit.score import PromptShieldScorer @@ -27,6 +30,21 @@ def promptshield_scorer() -> PromptShieldScorer: return PromptShieldScorer(prompt_shield_target=Mock()) +def generate_shield_response(response_text: str) -> Message: + return Message( + message_pieces=[ + MessagePiece( + role="assistant", + original_value=response_text, + original_value_data_type="text", + converted_value=response_text, + converted_value_data_type="text", + conversation_id=str(uuid.uuid4()), + ) + ] + ) + + @pytest.fixture def sample_delineated_prompt_as_str() -> str: sample: str = """ @@ -46,3 +64,21 @@ def test_prompt_shield_scorer_parsing_without_documents_analysis(promptshield_sc response_json_str = '{"userPromptAnalysis":{"attackDetected":false}}' result = promptshield_scorer._parse_response_to_boolean_list(response_json_str) assert result == [False, False] + + +async def test_prompt_shield_scorer_metadata_is_the_response_text(sqlite_instance, sample_response_json_str: str): + """The Score model only accepts flat str/int/float metadata values, so the endpoint + body has to be stored as the JSON text it arrived as rather than as a parsed object.""" + + target = MagicMock(spec=PromptTarget) + target.get_identifier.return_value = get_mock_target_identifier("MockShieldTarget") + target.send_prompt_async = AsyncMock(return_value=[generate_shield_response(sample_response_json_str)]) + + scorer = PromptShieldScorer(prompt_shield_target=target) + scores = await scorer.score_text_async(sample_response_json_str) + + assert len(scores) == 1 + # the sample body flags the document, not the user prompt + assert scores[0].get_value() is True + assert scores[0].score_metadata == {"response": sample_response_json_str} + assert json.loads(scores[0].score_metadata["response"]) == json.loads(sample_response_json_str)