Skip to content
Merged
14 changes: 10 additions & 4 deletions pyrit/score/float_scale/insecure_code_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,11 @@
from pyrit.score.float_scale.float_scale_scorer import MessageFloatScaleScorer
from pyrit.score.llm_scoring import _parse_judgment_observation, _run_llm_scoring_async
from pyrit.score.observation.execution import _ObservationEvidence
from pyrit.score.response_handler import JsonSchemaResponseHandler, ResponseHandler
from pyrit.score.response_handler import (
JsonSchemaResponseHandler,
NumericRangeResponseHandler,
ResponseHandler,
)
from pyrit.score.scorer_prompt_validator import ScorerPromptValidator
from pyrit.score.system_prompt import _render_system_prompt_template

Expand Down Expand Up @@ -118,9 +122,11 @@ def __init__(
# When the caller does not supply a response handler, the default JSON handler carries the
# schema (if any) declared by the system prompt and enforces the numeric score contract, so
# the round-trip forwards the schema to the scoring target. A caller-supplied handler owns
# its own response contract.
self._response_handler = response_handler or JsonSchemaResponseHandler(
response_schema=schema, numeric_value=True
# its own wire format.
wire_format_handler = response_handler or JsonSchemaResponseHandler(response_schema=schema, numeric_value=True)
# Keep score-domain validation in the parser callback so out-of-range values retry.
self._response_handler = NumericRangeResponseHandler(
response_handler=wire_format_handler, minimum_value=0, maximum_value=1
)

self._harm_categories = _normalize_harm_categories(harm_categories)
Expand Down
18 changes: 14 additions & 4 deletions pyrit/score/float_scale/self_ask_general_float_scale_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,11 @@
_parse_judgment_observation,
_run_llm_scoring_async,
)
from pyrit.score.response_handler import JsonSchemaResponseHandler, ResponseHandler
from pyrit.score.response_handler import (
Comment thread
hannahwestra25 marked this conversation as resolved.
JsonSchemaResponseHandler,
NumericRangeResponseHandler,
ResponseHandler,
)
from pyrit.score.scorer_prompt_validator import ScorerPromptValidator

if TYPE_CHECKING:
Expand Down Expand Up @@ -105,9 +109,9 @@ def __init__(
self._system_prompt_format_string = system_prompt_format_string
self._prompt_format_string = prompt_format_string
self._scale = scale
# A caller-supplied handler owns its own response contract; otherwise the default JSON
# handler carries the schema and enforces the numeric score contract for the round-trip.
self._response_handler = response_handler or JsonSchemaResponseHandler(
# A caller-supplied handler owns its own wire format; otherwise the default JSON handler
# carries the schema and enforces the numeric score contract for the round-trip.
wire_format_handler = response_handler or JsonSchemaResponseHandler(
score_value_output_key=score_value_output_key,
rationale_output_key=rationale_output_key,
description_output_key=description_output_key,
Expand All @@ -116,6 +120,12 @@ def __init__(
response_schema=response_json_schema,
numeric_value=True,
)
# Keep score-domain validation in the parser callback so out-of-range values retry.
self._response_handler = NumericRangeResponseHandler(
response_handler=wire_format_handler,
minimum_value=scale.minimum_value,
maximum_value=scale.maximum_value,
)

def _build_identifier(self) -> ComponentIdentifier:
"""
Expand Down
16 changes: 12 additions & 4 deletions pyrit/score/float_scale/self_ask_scale_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,11 @@
from pyrit.score.float_scale.numeric_scale import NumericRubric
from pyrit.score.llm_scoring import _parse_judgment_observation, _run_llm_scoring_async
from pyrit.score.observation.execution import _ObservationEvidence
from pyrit.score.response_handler import JsonSchemaResponseHandler, ResponseHandler
from pyrit.score.response_handler import (
JsonSchemaResponseHandler,
NumericRangeResponseHandler,
ResponseHandler,
)
from pyrit.score.scorer_prompt_validator import ScorerPromptValidator
from pyrit.score.system_prompt import _render_system_prompt_template

Expand Down Expand Up @@ -121,9 +125,13 @@ def __init__(
# When the caller does not supply a response handler, the default JSON handler carries the
# schema (if any) declared by the system prompt and enforces the numeric score contract, so
# the round-trip forwards the schema to the scoring target. A caller-supplied handler owns
# its own response contract.
self._response_handler = response_handler or JsonSchemaResponseHandler(
response_schema=schema, numeric_value=True
# its own wire format.
wire_format_handler = response_handler or JsonSchemaResponseHandler(response_schema=schema, numeric_value=True)
# Keep score-domain validation in the parser callback so out-of-range values retry.
self._response_handler = NumericRangeResponseHandler(
response_handler=wire_format_handler,
minimum_value=scale.minimum_value,
maximum_value=scale.maximum_value,
)

@classmethod
Expand Down
75 changes: 75 additions & 0 deletions pyrit/score/response_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -441,6 +441,81 @@ def parse(
return score


class NumericRangeResponseHandler(ResponseHandler):
"""Response-handler decorator that enforces a numeric score within ``[minimum_value, maximum_value]``."""

def __init__(self, *, response_handler: ResponseHandler, minimum_value: float, maximum_value: float) -> None:
"""
Initialize the decorator.

Args:
response_handler (ResponseHandler): Handler that parses the target's wire format.
minimum_value (float): The lowest accepted score value (inclusive).
maximum_value (float): The highest accepted score value (inclusive).
"""
self._response_handler = response_handler
self._minimum_value = minimum_value
self._maximum_value = maximum_value

@property
def json_response_config(self) -> JsonResponseConfig:
"""The wrapped handler's JSON-response request."""
return self._response_handler.json_response_config

def _replay_identifier(self) -> dict[str, Any] | None:
"""Return the wrapped parser identity with the accepted numeric range."""
wrapped = self._response_handler._get_replay_identifier()
if wrapped is None:
return None
return {
"handler": f"{type(self).__module__}.{type(self).__qualname__}",
"version": 1,
"wrapped": wrapped,
"minimum_value": self._minimum_value,
"maximum_value": self._maximum_value,
}

def parse(
self,
*,
response_text: str,
scorer_identifier: ComponentIdentifier,
scored_prompt_id: str | uuid.UUID,
category: Sequence[str] | str | None = None,
objective: str | None = None,
) -> UnvalidatedScore:
"""
Parse a response and require a finite numeric score within the configured range.

Returns:
UnvalidatedScore: The parsed score.

Raises:
InvalidJsonException: If the parsed value is not numeric, not finite, or out of range.
"""
score = self._response_handler.parse(
response_text=response_text,
scorer_identifier=scorer_identifier,
scored_prompt_id=scored_prompt_id,
category=category,
objective=objective,
)
try:
numeric_value = float(score.raw_score_value)
except ValueError:
raise InvalidJsonException(
message=f"score_value should be numeric, not {score.raw_score_value!r}."
) from None
if not math.isfinite(numeric_value) or not self._minimum_value <= numeric_value <= self._maximum_value:
raise InvalidJsonException(
message=(
f"score_value must be between {self._minimum_value} and {self._maximum_value}, "
f"not {score.raw_score_value!r}."
)
)
return score


class CallableResponseHandler(ResponseHandler):
"""
ResponseHandler that delegates parsing to a user-supplied callable.
Expand Down
23 changes: 23 additions & 0 deletions tests/unit/score/test_general_float_scale_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,6 +196,29 @@ async def test_general_float_scorer_score_async_min_max_range(patch_central_data
assert "Description." in score[0].score_value_description


async def test_general_float_scorer_retries_out_of_range_score(patch_central_database):
def _response(score_value: int) -> Message:
return Message(
message_pieces=[
MessagePiece(role="assistant", original_value=f'{{"score_value": {score_value}, "rationale": "r"}}')
]
)

chat_target = MagicMock()
chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget")
chat_target.send_prompt_async = AsyncMock(side_effect=[[_response(150)], [_response(40)]])

scorer = SelfAskGeneralFloatScaleScorer(
chat_target=chat_target,
system_prompt_format_string="Prompt.",
scale=DEFAULT_RANGE,
)
score = await scorer.score_text_async(text="prompt", objective="obj")

assert chat_target.send_prompt_async.call_count == 2
assert abs(float(score[0].score_value) - 0.4) < 1e-6


def test_general_float_scorer_init_invalid_min_max():
with pytest.raises(ValueError):
NumericRange(minimum_value=10, maximum_value=5, category="test")
Expand Down
18 changes: 18 additions & 0 deletions tests/unit/score/test_insecure_code_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,24 @@ async def test_insecure_code_scorer_real_response_handler_accepts_category_snaps
assert scores[0].get_value() == pytest.approx(0.5)


@pytest.mark.parametrize("out_of_range_value", ["-0.5", "1.5", "7"])
async def test_insecure_code_scorer_retries_out_of_range_score(mock_chat_target, out_of_range_value):
def _response(score_value: str) -> Message:
return Message(
message_pieces=[
MessagePiece(role="assistant", original_value=f'{{"score_value": {score_value}, "rationale": "r"}}')
]
)

mock_chat_target.send_prompt_async = AsyncMock(side_effect=[[_response(out_of_range_value)], [_response("0.3")]])
scorer = InsecureCodeScorer.from_harm_categories(chat_target=mock_chat_target)

scores = await scorer.score_text_async("sample code")

assert mock_chat_target.send_prompt_async.call_count == 2
assert scores[0].get_value() == pytest.approx(0.3)


async def test_score_async_unsupported_data_type_returns_empty(mock_chat_target, patch_central_database):
scorer = InsecureCodeScorer.from_harm_categories(chat_target=mock_chat_target)

Expand Down
65 changes: 65 additions & 0 deletions tests/unit/score/test_numeric_range_scorer_retries.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

"""Numeric-scale scorers retry an out-of-range judge score and give up once retries run out."""

from unittest.mock import AsyncMock, MagicMock

import pytest
from unit.mocks import get_mock_target_identifier

from pyrit.exceptions import InvalidJsonException
from pyrit.models import Message, MessagePiece
from pyrit.score import InsecureCodeScorer, NumericRange, NumericRubric, SelfAskScaleScorer
from pyrit.score.float_scale.self_ask_general_float_scale_scorer import SelfAskGeneralFloatScaleScorer


def _scale_scorer(chat_target):
return SelfAskScaleScorer.from_scale(
chat_target=chat_target,
scale=NumericRubric.from_yaml(SelfAskScaleScorer.ScalePaths.TREE_OF_ATTACKS_SCALE.value),
)


def _general_float_scorer(chat_target):
return SelfAskGeneralFloatScaleScorer(
chat_target=chat_target,
system_prompt_format_string="Prompt.",
scale=NumericRange(minimum_value=0, maximum_value=100, category="test"),
)


def _insecure_code_scorer(chat_target):
return InsecureCodeScorer.from_harm_categories(chat_target=chat_target)


@pytest.mark.parametrize(
("build_scorer", "out_of_range_value"),
[
(_scale_scorer, "11"),
(_general_float_scorer, "150"),
(_insecure_code_scorer, "1.5"),
],
ids=["SelfAskScaleScorer", "SelfAskGeneralFloatScaleScorer", "InsecureCodeScorer"],
)
async def test_out_of_range_score_on_every_attempt_raises_after_retries(
build_scorer, out_of_range_value: str, patch_central_database
):
response = Message(
message_pieces=[
MessagePiece(
role="assistant",
original_value=f'{{"score_value": "{out_of_range_value}", "rationale": "r", "description": "d"}}',
)
]
)
chat_target = MagicMock()
chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget")
chat_target.send_prompt_async = AsyncMock(return_value=[response])
scorer = build_scorer(chat_target)

with pytest.raises(InvalidJsonException):
await scorer.score_text_async(text="example text", objective="task")

# tests/unit/conftest.py pins RETRY_MAX_NUM_ATTEMPTS to 2.
assert chat_target.send_prompt_async.call_count == 2
62 changes: 61 additions & 1 deletion tests/unit/score/test_response_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,12 @@

from pyrit.exceptions import InvalidJsonException
from pyrit.models import ComponentIdentifier
from pyrit.score.response_handler import JsonSchemaResponseHandler, TrueFalseResponseHandler
from pyrit.score.response_handler import (
CallableResponseHandler,
JsonSchemaResponseHandler,
NumericRangeResponseHandler,
TrueFalseResponseHandler,
)

SCORER_IDENTIFIER = ComponentIdentifier(class_name="TestScorer", class_module=__name__)

Expand Down Expand Up @@ -87,3 +92,58 @@ def test_true_false_response_handler_rejects_value_outside_domain() -> None:
scorer_identifier=SCORER_IDENTIFIER,
scored_prompt_id="test-id",
)


@pytest.mark.parametrize("score_value", ["1", "5.5", "10"])
def test_numeric_range_response_handler_accepts_values_within_range(score_value: str) -> None:
handler = NumericRangeResponseHandler(
response_handler=JsonSchemaResponseHandler(numeric_value=True), minimum_value=1, maximum_value=10
)

score = handler.parse(
response_text=f'{{"score_value": "{score_value}", "rationale": "test"}}',
scorer_identifier=SCORER_IDENTIFIER,
scored_prompt_id="test-id",
)

assert score.raw_score_value == score_value


@pytest.mark.parametrize("score_value", ["0", "0.99", "10.01", "11", "-3"])
def test_numeric_range_response_handler_rejects_values_outside_range(score_value: str) -> None:
handler = NumericRangeResponseHandler(
response_handler=JsonSchemaResponseHandler(numeric_value=True), minimum_value=1, maximum_value=10
)

with pytest.raises(InvalidJsonException, match="must be between 1 and 10"):
handler.parse(
response_text=f'{{"score_value": "{score_value}", "rationale": "test"}}',
scorer_identifier=SCORER_IDENTIFIER,
scored_prompt_id="test-id",
)


@pytest.mark.parametrize("score_value", ["high", "nan"])
def test_numeric_range_response_handler_rejects_non_numeric_from_wrapped_handler(score_value: str) -> None:
# A caller-supplied wire-format handler may not validate numbers itself.
handler = NumericRangeResponseHandler(
response_handler=CallableResponseHandler(parser=lambda _: {"score_value": score_value, "rationale": "r"}),
minimum_value=1,
maximum_value=10,
)

with pytest.raises(InvalidJsonException):
handler.parse(response_text="ignored", scorer_identifier=SCORER_IDENTIFIER, scored_prompt_id="test-id")


def test_numeric_range_response_handler_replay_identifier_includes_range() -> None:
wrapped = JsonSchemaResponseHandler(numeric_value=True)
handler = NumericRangeResponseHandler(response_handler=wrapped, minimum_value=1, maximum_value=10)

identifier = handler._get_replay_identifier()

assert identifier is not None
assert identifier["wrapped"] == wrapped._get_replay_identifier()
assert identifier["minimum_value"] == 1
assert identifier["maximum_value"] == 10
assert handler.json_response_config == wrapped.json_response_config
Loading
Loading