Skip to content
Open
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
12 changes: 6 additions & 6 deletions pyrit/score/true_false/prompt_shield_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
42 changes: 39 additions & 3 deletions tests/unit/score/test_prompt_shield_scorer.py
Original file line number Diff line number Diff line change
@@ -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


Expand All @@ -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 = """
Expand All @@ -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)