From 0e48ed540896c7d911e2dc18a2edfe008199a777 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 13 Aug 2026 18:23:31 -0700 Subject: [PATCH 01/11] Add the scorable/expectation contract to Scorer (phase 1, code + tests) Give Scorer two inputs: a scorable (what to look at) and an expectation (what to look for). This lets a scorer answer questions that are not about a model response, which framework.md requires but the Message-only signature made impossible. This is a signature change, not a rewrite. A new MessageScorer intermediate base owns every message-shaped concern -- resolving a scorable to a Message, refusal and blocked-content substitution, piece validation, the role and error filters, the neutral fallback, and the exception wrapping around _score_async. The three direct Scorer subclasses re-parent onto it and no scorer body changes. - pyrit/models/score/ becomes a package holding every score type: the new scorables, ScoringExpectation and its conditions, ScoringScope, and the existing Score. - role_filter and skip_on_error_result move onto the message scorables, since both answer "which pieces count". - score_text_async and score_image_async keep their signatures and build a ContentScorable. Loose content stays ephemeral. - infer_objective_from_request is retired internally. ScorerEvaluator now reads each objective from the previous turn itself. - The legacy message-shaped parameters survive one release behind a DeprecationWarning, which the in-repo suite treats as an error. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyproject.toml | 5 + pyrit/executor/attack/multi_turn/crescendo.py | 31 +- .../executor/attack/multi_turn/red_teaming.py | 13 +- pyrit/models/__init__.py | 49 ++- pyrit/models/score/__init__.py | 44 ++ pyrit/models/score/expectation.py | 50 +++ pyrit/models/score/scorable.py | 121 ++++++ pyrit/models/{ => score}/score.py | 0 pyrit/models/score/scoring_scope.py | 28 ++ pyrit/score/__init__.py | 29 +- pyrit/score/audio_transcript_scorer.py | 7 +- pyrit/score/conversation_scorer.py | 3 +- pyrit/score/float_scale/float_scale_scorer.py | 13 +- pyrit/score/message_scorer.py | 261 ++++++++++++ pyrit/score/scorer.py | 331 +++++++-------- .../scorer_evaluation/scorer_evaluator.py | 24 +- .../float_scale_threshold_scorer.py | 18 +- .../true_false/true_false_composite_scorer.py | 15 +- .../true_false/true_false_inverter_scorer.py | 17 +- pyrit/score/true_false/true_false_scorer.py | 17 +- .../attack/multi_turn/test_crescendo.py | 6 +- .../multi_turn/test_crescendo_resilience.py | 15 +- .../attack/multi_turn/test_tree_of_attacks.py | 3 +- tests/unit/models/test_scorable.py | 144 +++++++ tests/unit/score/test_azure_content_filter.py | 14 +- .../score/test_conversation_history_scorer.py | 26 +- .../test_float_scale_threshold_scorer.py | 8 +- tests/unit/score/test_gandalf_scorer.py | 12 +- tests/unit/score/test_insecure_code_scorer.py | 16 +- tests/unit/score/test_message_scorer.py | 377 ++++++++++++++++++ tests/unit/score/test_plagiarism_scorer.py | 13 +- .../unit/score/test_question_answer_scorer.py | 12 +- tests/unit/score/test_scorer.py | 126 +++--- tests/unit/score/test_scorer_evaluator.py | 7 +- tests/unit/score/test_self_ask_category.py | 11 +- .../test_self_ask_question_answer_scorer.py | 6 +- tests/unit/score/test_self_ask_refusal.py | 3 +- tests/unit/score/test_shieldgemma_scorer.py | 6 +- tests/unit/score/test_substring.py | 4 +- .../score/test_true_false_composite_scorer.py | 27 +- tests/unit/score/test_true_false_inverter.py | 4 +- 41 files changed, 1469 insertions(+), 447 deletions(-) create mode 100644 pyrit/models/score/__init__.py create mode 100644 pyrit/models/score/expectation.py create mode 100644 pyrit/models/score/scorable.py rename pyrit/models/{ => score}/score.py (100%) create mode 100644 pyrit/models/score/scoring_scope.py create mode 100644 pyrit/score/message_scorer.py create mode 100644 tests/unit/models/test_scorable.py create mode 100644 tests/unit/score/test_message_scorer.py diff --git a/pyproject.toml b/pyproject.toml index 6baabf78d6..0cdc6cf6bb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -172,6 +172,11 @@ addopts = [ pythonpath = ["."] asyncio_default_fixture_loop_scope = "function" asyncio_mode = "auto" +filterwarnings = [ + # Keep the in-repo suite on the scorable/expectation contract. Tests that cover the + # shim itself opt back in with pytest.warns. + "error:Scorer\\.score_async:DeprecationWarning", +] [tool.ty] [tool.ty.rules] diff --git a/pyrit/executor/attack/multi_turn/crescendo.py b/pyrit/executor/attack/multi_turn/crescendo.py index 584065aa87..7b28bd773d 100644 --- a/pyrit/executor/attack/multi_turn/crescendo.py +++ b/pyrit/executor/attack/multi_turn/crescendo.py @@ -10,21 +10,11 @@ from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults from pyrit.common.path import EXECUTOR_SEED_PROMPT_PATH -from pyrit.exceptions import ( - ComponentRole, - execution_context, -) -from pyrit.executor.attack.component import ( - ConversationManager, - PrependedConversationConfig, -) +from pyrit.exceptions import ComponentRole, execution_context +from pyrit.executor.attack.component import ConversationManager, PrependedConversationConfig from pyrit.executor.attack.component.adversarial_conversation_manager import _AdversarialConversationManager from pyrit.executor.attack.component.modality_router import _ModalityFeedbackRouter -from pyrit.executor.attack.core import ( - AttackAdversarialConfig, - AttackConverterConfig, - AttackScoringConfig, -) +from pyrit.executor.attack.core import AttackAdversarialConfig, AttackConverterConfig, AttackScoringConfig from pyrit.executor.attack.multi_turn.multi_turn_attack_strategy import ( ConversationSession, MultiTurnAttackContext, @@ -41,18 +31,14 @@ ConversationType, Message, MessagePiece, + MessageScorable, Score, + ScoringExpectation, SeedPrompt, ) from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName, TargetRequirements -from pyrit.score import ( - FloatScaleThresholdScorer, - NumericRubric, - Scorer, - SelfAskRefusalScorer, - SelfAskScaleScorer, -) +from pyrit.score import FloatScaleThresholdScorer, NumericRubric, Scorer, SelfAskRefusalScorer, SelfAskScaleScorer from pyrit.score.score_utils import normalize_score_to_float if TYPE_CHECKING: @@ -680,9 +666,8 @@ async def _check_refusal_async(self, context: CrescendoAttackContext, objective: objective=context.objective, ): scores = await self._refusal_scorer.score_async( - message=context.last_response, - objective=objective, - skip_on_error_result=False, + scorable=MessageScorable(message=context.last_response), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) return scores[0] diff --git a/pyrit/executor/attack/multi_turn/red_teaming.py b/pyrit/executor/attack/multi_turn/red_teaming.py index 402eb1303b..59fc6627e8 100644 --- a/pyrit/executor/attack/multi_turn/red_teaming.py +++ b/pyrit/executor/attack/multi_turn/red_teaming.py @@ -18,11 +18,7 @@ get_adversarial_chat_messages, ) from pyrit.executor.attack.component.modality_router import _ModalityFeedbackRouter -from pyrit.executor.attack.core.attack_config import ( - AttackAdversarialConfig, - AttackConverterConfig, - AttackScoringConfig, -) +from pyrit.executor.attack.core.attack_config import AttackAdversarialConfig, AttackConverterConfig, AttackScoringConfig from pyrit.executor.attack.multi_turn.multi_turn_attack_strategy import ( ConversationSession, MultiTurnAttackContext, @@ -37,7 +33,9 @@ ConversationReference, ConversationType, Message, + MessageScorable, Score, + ScoringExpectation, ) from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName @@ -527,9 +525,8 @@ async def _score_response_async(self, *, context: MultiTurnAttackContext[Any]) - ): # score_async handles blocked, filtered, other errors scoring_results = await self._objective_scorer.score_async( - message=context.last_response, - role_filter="assistant", - objective=context.objective, + scorable=MessageScorable(message=context.last_response, role_filter="assistant"), + expectation=ScoringExpectation(objective=context.objective), ) objective_scores = scoring_results diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 289187f44c..5293afab4c 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -69,27 +69,34 @@ group_message_pieces_into_conversations, sort_message_pieces, ) -from pyrit.models.messages.chat_message import ( - ALLOWED_CHAT_MESSAGE_ROLES, - ChatMessage, - ChatMessagesDataset, - ToolCall, -) +from pyrit.models.messages.chat_message import ALLOWED_CHAT_MESSAGE_ROLES, ChatMessage, ChatMessagesDataset, ToolCall from pyrit.models.messages.conversation_reference import ConversationReference, ConversationType from pyrit.models.messages.conversation_retry import ConversationRetry, ConversationRetryReason -from pyrit.models.parameter import ( - ComponentType, - Parameter, - ParameterDestination, - RegistryReference, - display_choices, -) +from pyrit.models.parameter import ComponentType, Parameter, ParameterDestination, RegistryReference, display_choices from pyrit.models.question_answering import QuestionAnsweringDataset, QuestionAnsweringEntry, QuestionChoice from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT from pyrit.models.retry_event import RetryEvent -from pyrit.models.score import Score, ScoreType, UnvalidatedScore +from pyrit.models.score import ( + Condition, + ContentScorable, + ConversationScorable, + MessageReferenceScorable, + MessageScorable, + OutputMatches, + Scorable, + Score, + ScoreType, + ScoringExpectation, + ScoringScope, + SurfaceScorable, + ToolCalled, + ToolSequence, + TraceScorable, + UnvalidatedScore, + Volatility, +) # Seeds - import from new seeds submodule for forward compatibility # Also keep imports from old locations for backward compatibility @@ -144,11 +151,14 @@ "ComponentType", "compute_eval_hash", "config_hash", + "Condition", + "ContentScorable", "ConverterIdentifier", "Conversation", "ConversationReference", "ConversationRetry", "ConversationRetryReason", + "ConversationScorable", "ConversationStats", "ConversationType", "construct_response_from_request", @@ -181,9 +191,12 @@ "MEDIA_PATH_DATA_TYPES", "Message", "MessagePiece", + "MessageReferenceScorable", + "MessageScorable", "Modality", "NextMessageSystemPromptPaths", "ObjectiveTargetEvaluationIdentifier", + "OutputMatches", "Parameter", "ParameterDestination", "PromptDataType", @@ -194,8 +207,11 @@ "QuestionChoice", "REGISTRY_NAME_PATTERN", "ScaleDescription", + "Scorable", "Score", "ScoreType", + "ScoringExpectation", + "ScoringScope", "ScenarioEvaluationIdentifier", "ScorerEvaluationIdentifier", "ScorerIdentifier", @@ -218,6 +234,7 @@ "sort_message_pieces", "StrategyResult", "StrategyResultT", + "SurfaceScorable", "TARGET_EVAL_PARAM_FALLBACKS", "TARGET_EVAL_PARAMS", "TargetCapabilities", @@ -226,7 +243,11 @@ "TOKEN_USAGE_METADATA_PREFIX", "TokenUsage", "ToolCall", + "ToolCalled", + "ToolSequence", + "TraceScorable", "UnvalidatedScore", + "Volatility", "read_usage_int", "read_usage_value", "validate_registry_name", diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py new file mode 100644 index 0000000000..6bf0fdbe76 --- /dev/null +++ b/pyrit/models/score/__init__.py @@ -0,0 +1,44 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Score types: what is scored, what it is scored against, and the result. + +A scorer takes two inputs — a ``Scorable`` (what to look at) and a +``ScoringExpectation`` (what to look for) — and returns ``Score`` objects. +""" + +from pyrit.models.score.expectation import Condition, OutputMatches, ScoringExpectation, ToolCalled, ToolSequence +from pyrit.models.score.scorable import ( + ContentScorable, + ConversationScorable, + MessageReferenceScorable, + MessageScorable, + Scorable, + SurfaceScorable, + TraceScorable, + Volatility, +) +from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore +from pyrit.models.score.scoring_scope import ScoringScope + +__all__ = [ + "ComponentIdentifierField", + "Condition", + "ContentScorable", + "ConversationScorable", + "MessageReferenceScorable", + "MessageScorable", + "OutputMatches", + "Scorable", + "Score", + "ScoreType", + "ScoringExpectation", + "ScoringScope", + "SurfaceScorable", + "ToolCalled", + "ToolSequence", + "TraceScorable", + "UnvalidatedScore", + "Volatility", +] diff --git a/pyrit/models/score/expectation.py b/pyrit/models/score/expectation.py new file mode 100644 index 0000000000..0a89c3c080 --- /dev/null +++ b/pyrit/models/score/expectation.py @@ -0,0 +1,50 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from dataclasses import dataclass, field + + +@dataclass(frozen=True) +class ToolCalled: + """The named tool appears in the evidence.""" + + name: str + arguments: dict[str, str] | None = None + + +@dataclass(frozen=True) +class ToolSequence: + """The named tools appear, in this relative order.""" + + names: tuple[str, ...] + + +@dataclass(frozen=True) +class OutputMatches: + """The response contains or equals a ground-truth answer.""" + + value: str + + +Condition = ToolCalled | ToolSequence | OutputMatches + + +@dataclass(frozen=True) +class ScoringExpectation: + """ + What a scorer scores against. + + An expectation is a single parameter that attacks forward without inspecting it, + so a question authored in a technique configuration or a seed can reach a scorer + through an attack that knows nothing about it. + + Conditions are neutral about polarity: a condition says what to detect, never + whether detecting it is good or bad. Wrap a scorer in ``TrueFalseInverterScorer`` + to express the negative case. + """ + + objective: str | None = None + conditions: tuple[Condition, ...] = () + extra: dict[str, str] = field(default_factory=dict) diff --git a/pyrit/models/score/scorable.py b/pyrit/models/score/scorable.py new file mode 100644 index 0000000000..5864a22840 --- /dev/null +++ b/pyrit/models/score/scorable.py @@ -0,0 +1,121 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum +from typing import TYPE_CHECKING, ClassVar + +if TYPE_CHECKING: + import uuid + + from pyrit.models.literals import ChatMessageRole, PromptDataType + from pyrit.models.messages import Message + from pyrit.models.score.scoring_scope import ScoringScope + + +class Volatility(Enum): + """Whether two reads of the same scorable can disagree.""" + + STABLE = "stable" + APPEND_ONLY = "append_only" + MUTABLE = "mutable" + + +@dataclass(frozen=True) +class ConversationScorable: + """ + A stored conversation, resolved from memory. + + A conversation is append-only rather than stable because it grows while the attack + runs; reading it after turn one and after turn five gives different evidence. Leave + ``through_sequence`` unset to mean "as far as it goes right now"; set it to name a + reproducible prefix. + """ + + volatility: ClassVar[Volatility] = Volatility.APPEND_ONLY + + conversation_id: str + through_sequence: int | None = None + role_filter: ChatMessageRole | None = None + + +@dataclass(frozen=True) +class MessageScorable: + """ + A message the caller already holds. + + Use this when the message is in hand — an attack scoring the response it just + received, or a caller scoring a message whose pieces were never persisted. Use + ``MessageReferenceScorable`` when only the piece ids are known. + """ + + volatility: ClassVar[Volatility] = Volatility.STABLE + + message: Message + role_filter: ChatMessageRole | None = None + skip_on_error_result: bool = False + + +@dataclass(frozen=True) +class MessageReferenceScorable: + """ + Specific message pieces, resolved from memory by id. + + Use this when the pieces are persisted and only their ids are known. Use + ``MessageScorable`` when the caller already holds the message. + """ + + volatility: ClassVar[Volatility] = Volatility.STABLE + + message_piece_ids: tuple[uuid.UUID | str, ...] + role_filter: ChatMessageRole | None = None + skip_on_error_result: bool = False + + +@dataclass(frozen=True) +class ContentScorable: + """Loose content with no conversation behind it.""" + + volatility: ClassVar[Volatility] = Volatility.STABLE + + value: str + data_type: PromptDataType = "text" + + +@dataclass(frozen=True) +class SurfaceScorable: + """A location that may or may not have been written.""" + + volatility: ClassVar[Volatility] = Volatility.MUTABLE + + uri: str + surface: str = "file" + scope: ScoringScope | None = None + + +@dataclass(frozen=True) +class TraceScorable: + """ + A recorded agent run, or named steps within it. + + A span is the tracing term for one recorded step — a single tool call, model + request, or nested operation. A trace is the tree of spans for one run. + """ + + volatility: ClassVar[Volatility] = Volatility.APPEND_ONLY + + trace_ids: tuple[str, ...] = () + span_name: str | None = None + scope: ScoringScope | None = None + + +Scorable = ( + ConversationScorable + | MessageScorable + | MessageReferenceScorable + | ContentScorable + | SurfaceScorable + | TraceScorable +) diff --git a/pyrit/models/score.py b/pyrit/models/score/score.py similarity index 100% rename from pyrit/models/score.py rename to pyrit/models/score/score.py diff --git a/pyrit/models/score/scoring_scope.py b/pyrit/models/score/scoring_scope.py new file mode 100644 index 0000000000..8247c5612f --- /dev/null +++ b/pyrit/models/score/scoring_scope.py @@ -0,0 +1,28 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from datetime import datetime + + +@dataclass(frozen=True) +class ScoringScope: + """ + Bounds a scorable that names a location rather than a stored record. + + A surface URI or a trace id names a place, not a thing, so it cannot say on its + own which run produced what is there now. A scope narrows that question with a + time window and correlation labels. + + Labels are the extension point for frameworks built on PyRIT that control how + evidence is emitted and can therefore supply stronger correlation keys. PyRIT + never interprets them. + """ + + window: tuple[datetime, datetime] | None = None + labels: dict[str, str] = field(default_factory=dict) diff --git a/pyrit/score/__init__.py b/pyrit/score/__init__.py index 0647026606..24f449e27c 100644 --- a/pyrit/score/__init__.py +++ b/pyrit/score/__init__.py @@ -19,10 +19,7 @@ FloatScaleScorerByCategory, ) from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer -from pyrit.score.float_scale.insecure_code_scorer import ( - InsecureCodeScorer, - render_insecure_code_system_prompt, -) +from pyrit.score.float_scale.insecure_code_scorer import InsecureCodeScorer, render_insecure_code_system_prompt from pyrit.score.float_scale.likert_scale import LikertScale, LikertScaleEntry from pyrit.score.float_scale.numeric_scale import NumericRange, NumericRubric from pyrit.score.float_scale.plagiarism_scorer import PlagiarismMetric, PlagiarismScorer @@ -33,15 +30,9 @@ SelfAskLikertScorer, render_likert_system_prompt, ) -from pyrit.score.float_scale.self_ask_scale_scorer import ( - SelfAskScaleScorer, - render_scale_system_prompt, -) -from pyrit.score.response_handler import ( - CallableResponseHandler, - JsonSchemaResponseHandler, - ResponseHandler, -) +from pyrit.score.float_scale.self_ask_scale_scorer import SelfAskScaleScorer, render_scale_system_prompt +from pyrit.score.message_scorer import MessageScorer +from pyrit.score.response_handler import CallableResponseHandler, JsonSchemaResponseHandler, ResponseHandler from pyrit.score.scorer import Scorer from pyrit.score.scorer_evaluation.metrics_type import MetricsType, RegistryUpdateBehavior from pyrit.score.scorer_evaluation.scorer_metrics import ( @@ -62,11 +53,7 @@ from pyrit.score.true_false.gandalf_scorer import GandalfScorer from pyrit.score.true_false.llamaguard_parser import LLAMAGUARD_3_CATEGORY_CODES, parse_llamaguard_response from pyrit.score.true_false.llamaguard_policy import LlamaGuardCategory, LlamaGuardPolicy -from pyrit.score.true_false.llamaguard_scorer import ( - LlamaGuardMessageRole, - LlamaGuardScorer, - render_llamaguard_prompt, -) +from pyrit.score.true_false.llamaguard_scorer import LlamaGuardMessageRole, LlamaGuardScorer, render_llamaguard_prompt from pyrit.score.true_false.prompt_shield_scorer import PromptShieldScorer from pyrit.score.true_false.question_answer_scorer import QuestionAnswerScorer from pyrit.score.true_false.regex.anthrax_keyword_scorer import AnthraxKeywordScorer @@ -109,10 +96,7 @@ ShieldGemmaMessageRole, ShieldGemmaPolicy, ) -from pyrit.score.true_false.shieldgemma_scorer import ( - ShieldGemmaScorer, - render_shieldgemma_prompt, -) +from pyrit.score.true_false.shieldgemma_scorer import ShieldGemmaScorer, render_shieldgemma_prompt from pyrit.score.true_false.substring_scorer import SubStringScorer from pyrit.score.true_false.true_false_composite_scorer import TrueFalseCompositeScorer from pyrit.score.true_false.true_false_inverter_scorer import TrueFalseInverterScorer @@ -204,6 +188,7 @@ def __getattr__(name: str) -> object: "LlamaGuardPolicy", "LlamaGuardScorer", "MarkdownInjectionScorer", + "MessageScorer", "MethKeywordScorer", "MetricsType", "NerveAgentKeywordScorer", diff --git a/pyrit/score/audio_transcript_scorer.py b/pyrit/score/audio_transcript_scorer.py index b49d32842e..1977f52372 100644 --- a/pyrit/score/audio_transcript_scorer.py +++ b/pyrit/score/audio_transcript_scorer.py @@ -11,7 +11,7 @@ from pyrit.converter import AzureSpeechAudioToTextConverter from pyrit.memory import CentralMemory -from pyrit.models import MessagePiece, Score +from pyrit.models import MessagePiece, MessageScorable, Score, ScoringExpectation from pyrit.score.scorer import Scorer logger = logging.getLogger(__name__) @@ -185,7 +185,10 @@ async def _score_audio_async(self, *, message_piece: MessagePiece, objective: st memory.add_message_to_memory(request=text_message) # Score the transcript - transcript_scores = await self.text_scorer.score_async(message=text_message, objective=objective) + transcript_scores = await self.text_scorer.score_async( + scorable=MessageScorable(message=text_message), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, + ) # Add context to indicate this was scored from audio transcription for score in transcript_scores: diff --git a/pyrit/score/conversation_scorer.py b/pyrit/score/conversation_scorer.py index 81acb43624..1a78e88793 100644 --- a/pyrit/score/conversation_scorer.py +++ b/pyrit/score/conversation_scorer.py @@ -6,6 +6,7 @@ from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer +from pyrit.score.message_scorer import MessageScorer from pyrit.score.scorer import Scorer from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -14,7 +15,7 @@ from uuid import UUID -class ConversationScorer(Scorer, ABC): +class ConversationScorer(MessageScorer, ABC): """ Scorer that evaluates entire conversation history rather than individual messages. diff --git a/pyrit/score/float_scale/float_scale_scorer.py b/pyrit/score/float_scale/float_scale_scorer.py index 538ea44d1c..6206dadeea 100644 --- a/pyrit/score/float_scale/float_scale_scorer.py +++ b/pyrit/score/float_scale/float_scale_scorer.py @@ -5,11 +5,8 @@ from typing import TYPE_CHECKING -from pyrit.models import ( - Message, - Score, -) -from pyrit.score.scorer import Scorer +from pyrit.models import Message, Score +from pyrit.score.message_scorer import MessageScorer if TYPE_CHECKING: from pyrit.prompt_target.common.prompt_target import PromptTarget @@ -17,7 +14,7 @@ from pyrit.score.scorer_prompt_validator import ScorerPromptValidator -class FloatScaleScorer(Scorer): +class FloatScaleScorer(MessageScorer): """ Base class for scorers that return floating-point scores in the range [0, 1]. @@ -127,9 +124,7 @@ def get_scorer_metrics(self) -> HarmScorerMetrics | None: Returns: HarmScorerMetrics: The metrics for this scorer, or None if not found or not configured. """ - from pyrit.score.scorer_evaluation.scorer_metrics_io import ( - find_harm_metrics_by_eval_hash, - ) + from pyrit.score.scorer_evaluation.scorer_metrics_io import find_harm_metrics_by_eval_hash if self.evaluation_file_mapping is None or self.evaluation_file_mapping.harm_category is None: return None diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py new file mode 100644 index 0000000000..80787a45fb --- /dev/null +++ b/pyrit/score/message_scorer.py @@ -0,0 +1,261 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException +from pyrit.models import ( + ContentScorable, + Message, + MessagePiece, + MessageReferenceScorable, + MessageScorable, + Scorable, + Score, + ScoringExpectation, + group_message_pieces_into_conversations, +) +from pyrit.score.scorer import Scorer + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + +logger = logging.getLogger(__name__) + +#: Scorable kinds a MessageScorer can reduce to a single Message. +_SUPPORTED_SCORABLES = (MessageScorable, MessageReferenceScorable, ContentScorable) + + +def extract_objective_from_previous_turn(*, message: Message, memory: MemoryInterface) -> str: + """ + Read the text of the turn before an assistant message and use it as the objective. + + This is what to look for, so it belongs to the caller that builds the expectation, not + to the scorer. It lives here because it is message-shaped. + + Args: + message (Message): The assistant message whose previous turn supplies the objective. + memory (MemoryInterface): Memory holding the conversation. + + Returns: + str: The previous turn's text, or an empty string when there is none. + """ + if not message.message_pieces: + return "" + + piece = message.get_piece() + + if piece.api_role != "assistant": + return "" + + conversation = memory.get_message_pieces(conversation_id=piece.conversation_id) + if not conversation: + return "" + + last_prompt = max(conversation, key=lambda x: x.sequence) + + return "\n".join( + [ + piece.original_value + for piece in conversation + if piece.sequence == last_prompt.sequence - 1 and piece.original_value_data_type == "text" + ] + ) + + +class MessageScorer(Scorer): + """ + Base class for scorers whose evidence is a single message. + + Every message-shaped concern lives here: resolving a message scorable to a ``Message``, + substituting refusal and blocked content, validating pieces, applying the role and error + filters, and falling back to a neutral score. ``Scorer`` stays agnostic about what a + scorable is, so scorers over other kinds of evidence can sit beside this one. + + Subclasses implement ``_score_async``, which still receives a ``Message``. + """ + + async def _score_scorable_async( + self, + *, + scorable: Scorable, + expectation: ScoringExpectation | None, + infer_objective_from_request: bool = False, + ) -> list[Score]: + """ + Resolve a message scorable and score the message it names. + + Args: + scorable (Scorable): A ``MessageScorable``, ``MessageReferenceScorable``, or + ``ContentScorable``. + expectation (ScoringExpectation | None): What to look for. + infer_objective_from_request (bool): Deprecated; read the objective from the + previous turn when the expectation carries none. + + Returns: + list[Score]: The scores, or an empty list when a filter skipped the message. + + Raises: + TypeError: If the scorable is not message-shaped. + ScorerLLMResponseBlockedException: If the scorer's own LLM response is blocked by + content filtering and ``raise_if_scorer_blocks`` is True (the default). + PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). + RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). + """ + if not isinstance(scorable, _SUPPORTED_SCORABLES): + raise TypeError( + f"{self.__class__.__name__} scores messages, so it cannot score {type(scorable).__name__}. " + "Pass a MessageScorable, a MessageReferenceScorable, or a ContentScorable." + ) + + message = self._resolve_message(scorable) + role_filter = getattr(scorable, "role_filter", None) + skip_on_error_result = getattr(scorable, "skip_on_error_result", False) + objective = expectation.objective if expectation else None + + # Structured refusals are persisted as blocked error pieces, but scorers should + # receive the refusal explanation as text. Keep response_error="blocked" so + # refusal scorers can still use their deterministic blocked-response path. + scoring_message = self._apply_structured_refusal_substitution(message) + + # When score_blocked_content is enabled, blocked pieces with partial content + # take precedence and are replaced with text substitutes (response_error="none"). + if self.score_blocked_content: + scoring_message = self._apply_blocked_content_substitution(scoring_message) + + self._validator.validate(scoring_message, objective=objective) + + if role_filter is not None and message.get_piece().role != role_filter: + logger.debug("Skipping scoring due to role filter mismatch.") + return [] + + if skip_on_error_result and self._should_skip_on_error(message): + return [] + + if infer_objective_from_request and (not objective): + objective = extract_objective_from_previous_turn(message=message, memory=self._memory) + + try: + scores = await self._score_async( + scoring_message, + objective=objective, + ) + except ScorerLLMResponseBlockedException as e: + # The scorer's own LLM response was content-filtered. By default this is a real + # error and propagates; when raise_if_scorer_blocks is False, fall back to the + # scorer's type default (False / 0.0) instead. The decision lives here in the + # scorer, not the transport (see doc/code/framework.md). + if self.raise_if_scorer_blocks: + e.message = f"Error in scorer {self.__class__.__name__}: {e.message}" + e.args = (f"Status Code: {e.status_code}, Message: {e.message}",) + raise + logger.info( + "Scorer %s LLM response was blocked by content filtering; " + "returning default score (raise_if_scorer_blocks=False).", + self.__class__.__name__, + ) + scores = self._build_fallback_score( + message=scoring_message, + objective=objective, + scorer_response_blocked=True, + ) + except PyritException as e: + # Re-raise PyRIT exceptions with enhanced context while preserving type for retry decorators + e.message = f"Error in scorer {self.__class__.__name__}: {e.message}" + e.args = (f"Status Code: {e.status_code}, Message: {e.message}",) + raise + except Exception as e: + # Wrap non-PyRIT exceptions for better error tracing + raise RuntimeError(f"Error in scorer {self.__class__.__name__}: {str(e)}") from e + + if not scores and scoring_message.message_pieces: + scores = self._build_fallback_score(message=scoring_message, objective=objective) + + self._drop_ephemeral_score_links(message=scoring_message, scores=scores) + + return scores + + def _resolve_message(self, scorable: MessageScorable | MessageReferenceScorable | ContentScorable) -> Message: + """ + Return the message a message-shaped scorable names. + + Loose content has no conversation behind it, so it becomes a message that is marked + as never persisted. Phase 2 gives scores a scorable of their own and removes this. + + Returns: + Message: The message to score. + + Raises: + ValueError: If the referenced pieces are not in memory or do not form one message. + """ + if isinstance(scorable, MessageScorable): + return scorable.message + + if isinstance(scorable, ContentScorable): + piece = MessagePiece( + role="user", + original_value=scorable.value, + original_value_data_type=scorable.data_type, + ) + piece.not_in_memory = True + return Message(message_pieces=[piece]) + + pieces = self._memory.get_message_pieces(prompt_ids=list(scorable.message_piece_ids)) + if not pieces: + raise ValueError(f"No message pieces found in memory for ids {list(scorable.message_piece_ids)}.") + + conversations = group_message_pieces_into_conversations(pieces) + messages = [message for conversation in conversations for message in conversation] + if len(messages) != 1: + raise ValueError( + f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " + "Reference pieces from a single message, or use a ConversationScorable." + ) + return messages[0] + + def _should_skip_on_error(self, message: Message) -> bool: + """ + Return whether an errored message should be skipped rather than scored. + + Returns: + bool: True when the message should not be scored. + """ + if not message.is_error(): + return False + + error_pieces = [ + piece for piece in message.message_pieces if piece.has_error() or piece.converted_value_data_type == "error" + ] + # SDK-provided structured refusals stay scoreable: the refusal text is the evidence. + only_structured_refusals = all(piece.structured_refusal is not None for piece in error_pieces) + # When score_blocked_content is enabled and the message has partial content, + # don't skip — let _score_async handle the substitution. + all_errors_have_partial_content = all( + piece.is_blocked() and piece.prompt_metadata.get("partial_content") for piece in error_pieces + ) + if only_structured_refusals or (self.score_blocked_content and all_errors_have_partial_content): + return False + + logger.debug("Skipping scoring due to error in message and skip_on_error=True.") + return True + + @staticmethod + def _drop_ephemeral_score_links(*, message: Message, scores: list[Score]) -> None: + """ + Clear the piece link on scores that point at pieces which were never persisted. + + Memory cannot link a score to a piece it never stored, but the score itself is + still worth keeping. + """ + ephemeral_piece_ids = { + piece.id for piece in message.message_pieces if piece.not_in_memory and piece.id is not None + } + if not ephemeral_piece_ids: + return + + for score in scores: + if score.message_piece_id in ephemeral_piece_ids: + score.message_piece_id = None # type: ignore[ty:invalid-assignment] diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index bd3a85c686..1a1c07df9c 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -7,26 +7,25 @@ import asyncio import logging from abc import abstractmethod -from typing import ( - TYPE_CHECKING, - Any, - ClassVar, - cast, -) +from typing import TYPE_CHECKING, Any, ClassVar, cast -from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException +from pyrit.common.deprecation import print_deprecation_message from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ( ChatMessageRole, ComponentIdentifier, + ContentScorable, Identifiable, Message, MessagePiece, + MessageScorable, PromptResponseError, + Scorable, Score, ScorerEvaluationIdentifier, ScorerIdentifier, ScoreType, + ScoringExpectation, ) from pyrit.prompt_target.batch_helper import batch_task_async from pyrit.prompt_target.common.target_requirements import TargetRequirements @@ -36,14 +35,15 @@ from pyrit.prompt_target import PromptTarget from pyrit.score.scorer_evaluation.metrics_type import RegistryUpdateBehavior - from pyrit.score.scorer_evaluation.scorer_evaluator import ( - ScorerEvalDatasetFiles, - ) + from pyrit.score.scorer_evaluation.scorer_evaluator import ScorerEvalDatasetFiles from pyrit.score.scorer_evaluation.scorer_metrics import ScorerMetrics from pyrit.score.scorer_prompt_validator import ScorerPromptValidator logger = logging.getLogger(__name__) +#: Release in which the message-shaped ``score_async`` parameters are removed. +LEGACY_SCORE_ASYNC_REMOVED_IN = "1.3.0" + class Scorer(Identifiable, abc.ABC): """ @@ -205,125 +205,158 @@ def _create_identifier( async def score_async( self, - message: Message, + message: Message | None = None, *, + scorable: Scorable | None = None, + expectation: ScoringExpectation | None = None, objective: str | None = None, role_filter: ChatMessageRole | None = None, skip_on_error_result: bool = False, infer_objective_from_request: bool = False, ) -> list[Score]: """ - Score the message, add the results to the database, and return a list of Score objects. + Score a scorable against an expectation, persist the results, and return them. + + A scorer takes two inputs: the ``scorable`` says what to look at, and the + ``expectation`` says what to look for. Keeping them separate is what lets an attack + forward a question it does not understand to a scorer that does. + + This signature is experimental for one release and may change. Args: - message (Message): The message to be scored. - objective (str | None): The task or objective based on which the message should be scored. - Defaults to None. - role_filter (ChatMessageRole | None): Only score messages with this exact stored role. - Use "assistant" to score only real assistant responses, or "simulated_assistant" - to score only simulated responses. Defaults to None (no filtering). - skip_on_error_result (bool): If True, skip scoring if the message contains an error. - SDK-provided structured refusals remain scoreable. When self.score_blocked_content - is also True, blocked responses with partial content will still be scored instead - of skipping. Defaults to False. - infer_objective_from_request (bool): If True, infer the objective from the message's previous request - when objective is not provided. Defaults to False. + message (Message | None): Deprecated. Pass + ``scorable=MessageScorable(message=...)`` instead. + scorable (Scorable | None): What to look at. + expectation (ScoringExpectation | None): What to look for. Defaults to None. + objective (str | None): Deprecated. Pass + ``expectation=ScoringExpectation(objective=...)`` instead. + role_filter (ChatMessageRole | None): Deprecated. Set ``role_filter`` on the + message scorable instead. + skip_on_error_result (bool): Deprecated. Set ``skip_on_error_result`` on the + message scorable instead. + infer_objective_from_request (bool): Deprecated. Resolve the objective at the + call site and pass it on the expectation instead. Returns: list[Score]: A list of Score objects representing the results. Raises: - ScorerLLMResponseBlockedException: If the scorer's own LLM response is blocked by - content filtering and ``raise_if_scorer_blocks`` is True (the default). - PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). - RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). - """ - # Structured refusals are persisted as blocked error pieces, but scorers should - # receive the refusal explanation as text. Keep response_error="blocked" so - # refusal scorers can still use their deterministic blocked-response path. - scoring_message = self._apply_structured_refusal_substitution(message) - - # When score_blocked_content is enabled, blocked pieces with partial content - # take precedence and are replaced with text substitutes (response_error="none"). - if self.score_blocked_content: - scoring_message = self._apply_blocked_content_substitution(scoring_message) + ValueError: If the scorable inputs are missing, duplicated, or combined with + parameters that do not apply to them. + TypeError: If this scorer does not support this kind of scorable. + """ + resolved_scorable, resolved_expectation = self._resolve_score_inputs( + message=message, + scorable=scorable, + expectation=expectation, + objective=objective, + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + infer_objective_from_request=infer_objective_from_request, + ) - self._validator.validate(scoring_message, objective=objective) + scores = await self._score_scorable_async( + scorable=resolved_scorable, + expectation=resolved_expectation, + infer_objective_from_request=infer_objective_from_request, + ) - if role_filter is not None and message.get_piece().role != role_filter: - logger.debug("Skipping scoring due to role filter mismatch.") + if not scores: return [] - if skip_on_error_result and message.is_error(): - error_pieces = [ - piece - for piece in message.message_pieces - if piece.has_error() or piece.converted_value_data_type == "error" - ] - only_structured_refusals = all(piece.structured_refusal is not None for piece in error_pieces) - # When score_blocked_content is enabled and the message has partial content, - # don't skip — let _score_async handle the substitution. - all_errors_have_partial_content = all( - piece.is_blocked() and piece.prompt_metadata.get("partial_content") for piece in error_pieces - ) - if not only_structured_refusals and not (self.score_blocked_content and all_errors_have_partial_content): - logger.debug("Skipping scoring due to error in message and skip_on_error=True.") - return [] + self.validate_return_scores(scores=scores) + self._memory.add_scores_to_memory(scores=scores) - if infer_objective_from_request and (not objective): - objective = self._extract_objective_from_response(message) + return scores - try: - scores = await self._score_async( - scoring_message, - objective=objective, + def _resolve_score_inputs( + self, + *, + message: Message | None, + scorable: Scorable | None, + expectation: ScoringExpectation | None, + objective: str | None, + role_filter: ChatMessageRole | None, + skip_on_error_result: bool, + infer_objective_from_request: bool, + ) -> tuple[Scorable, ScoringExpectation | None]: + """ + Map the deprecated message-shaped parameters onto a scorable and an expectation. + + Returns: + tuple[Scorable, ScoringExpectation | None]: The resolved scoring inputs. + + Raises: + ValueError: If the inputs are missing, duplicated, or combined with parameters + that do not apply to them. + """ + if message is not None and scorable is not None: + raise ValueError("Pass either 'message' or 'scorable', not both.") + if message is None and scorable is None: + raise ValueError("Either 'message' or 'scorable' must be provided.") + if objective is not None and expectation is not None: + raise ValueError("Pass either 'objective' or 'expectation', not both.") + if scorable is not None and (role_filter is not None or skip_on_error_result): + raise ValueError( + "'role_filter' and 'skip_on_error_result' are fields on the message scorable. " + "Set them on the scorable instead of passing them to score_async." ) - except ScorerLLMResponseBlockedException as e: - # The scorer's own LLM response was content-filtered. By default this is a real - # error and re-raised; when raise_if_scorer_blocks is False, fall back to the - # scorer's type default (False / 0.0) instead. The decision lives here in the - # Scorer, not the transport (see doc/code/framework.md). - if self.raise_if_scorer_blocks: - e.message = f"Error in scorer {self.__class__.__name__}: {e.message}" - e.args = (f"Status Code: {e.status_code}, Message: {e.message}",) - raise - logger.info( - "Scorer %s LLM response was blocked by content filtering; " - "returning default score (raise_if_scorer_blocks=False).", - self.__class__.__name__, + + uses_legacy_parameters = ( + message is not None + or objective is not None + or role_filter is not None + or skip_on_error_result + or infer_objective_from_request + ) + if uses_legacy_parameters: + print_deprecation_message( + old_item="Scorer.score_async(message=..., objective=..., role_filter=..., " + "skip_on_error_result=..., infer_objective_from_request=...)", + new_item="Scorer.score_async(scorable=..., expectation=...)", + removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, ) - scores = self._build_fallback_score( - message=scoring_message, - objective=objective, - scorer_response_blocked=True, + + if scorable is None: + # A supplied message maps to MessageScorable, never ConversationScorable: the + # caller asked about this message, not about everything stored alongside it. + scorable = MessageScorable( + message=cast("Message", message), + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, ) - except PyritException as e: - # Re-raise PyRIT exceptions with enhanced context while preserving type for retry decorators - e.message = f"Error in scorer {self.__class__.__name__}: {e.message}" - e.args = (f"Status Code: {e.status_code}, Message: {e.message}",) - raise - except Exception as e: - # Wrap non-PyRIT exceptions for better error tracing - raise RuntimeError(f"Error in scorer {self.__class__.__name__}: {str(e)}") from e - - if not scores and scoring_message.message_pieces: - scores = self._build_fallback_score(message=scoring_message, objective=objective) - self.validate_return_scores(scores=scores) + if objective is not None: + expectation = ScoringExpectation(objective=objective) - # For pieces flagged not-in-memory, drop the FK on any score that points at them - # so memory doesn't try to link a score to a piece that was never persisted. - ephemeral_piece_ids = { - piece.id for piece in scoring_message.message_pieces if piece.not_in_memory and piece.id is not None - } - if ephemeral_piece_ids: - for score in scores: - if score.message_piece_id in ephemeral_piece_ids: - score.message_piece_id = None # type: ignore[ty:invalid-assignment] + return scorable, expectation - self._memory.add_scores_to_memory(scores=scores) + async def _score_scorable_async( + self, + *, + scorable: Scorable, + expectation: ScoringExpectation | None, + infer_objective_from_request: bool = False, + ) -> list[Score]: + """ + Score a scorable this scorer supports. - return scores + Subclasses implement this for the scorable kinds they handle and raise + ``TypeError`` for the rest. ``MessageScorer`` handles the message-shaped kinds. + + An implementation returns an empty list when a filter skipped the scorable without + scoring it. An empty list bypasses ``validate_return_scores`` and persistence. + + Args: + scorable (Scorable): What to look at. + expectation (ScoringExpectation | None): What to look for. + infer_objective_from_request (bool): Deprecated; resolve the objective from the + stored conversation when the expectation carries none. + + Raises: + TypeError: If the scorer does not support this kind of scorable. + """ + raise TypeError(f"{self.__class__.__name__} does not support scorable {type(scorable).__name__}.") async def _score_async(self, message: Message, *, objective: str | None = None) -> list[Score]: """ @@ -619,18 +652,11 @@ async def score_text_async(self, text: str, *, objective: str | None = None) -> Returns: list[Score]: A list of Score objects representing the results. """ - request = Message( - message_pieces=[ - MessagePiece( - role="user", - original_value=text, - ) - ] + return await self.score_async( + scorable=ContentScorable(value=text), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) - request.message_pieces[0].not_in_memory = True - return await self.score_async(request, objective=objective) - async def score_image_async(self, image_path: str, *, objective: str | None = None) -> list[Score]: """ Score the given image using the chat target. @@ -642,19 +668,11 @@ async def score_image_async(self, image_path: str, *, objective: str | None = No Returns: list[Score]: A list of Score objects representing the results. """ - request = Message( - message_pieces=[ - MessagePiece( - role="user", - original_value=image_path, - original_value_data_type="image_path", - ) - ] + return await self.score_async( + scorable=ContentScorable(value=image_path, data_type="image_path"), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) - request.message_pieces[0].not_in_memory = True - return await self.score_async(request, objective=objective) - async def score_prompts_batch_async( self, *, @@ -694,16 +712,20 @@ async def score_prompts_batch_async( if len(messages) == 0: return [] + scorables = [ + MessageScorable(message=message, role_filter=role_filter, skip_on_error_result=skip_on_error_result) + for message in messages + ] + expectations = [ScoringExpectation(objective=objective) for objective in objectives] + # Some scorers do not have an associated prompt target; batch helper validates RPM only when present prompt_target = getattr(self, "_prompt_target", None) results = await batch_task_async( task_func=self.score_async, - task_arguments=["message", "objective"], + task_arguments=["scorable", "expectation"], prompt_target=cast("PromptTarget", prompt_target), batch_size=batch_size, - items_to_batch=[messages, objectives], - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, + items_to_batch=[scorables, expectations], infer_objective_from_request=infer_objective_from_request, ) @@ -764,7 +786,9 @@ def scale_value_float(self, value: float, min_value: float, max_value: float) -> def _extract_objective_from_response(self, response: Message) -> str: """ - Extract an objective from the response using the last request (if it exists). + Read the objective from the turn before an assistant response. + + Deprecated: use ``pyrit.score.message_scorer.extract_objective_from_previous_turn``. Args: response (Message): The response to extract the objective from. @@ -772,25 +796,14 @@ def _extract_objective_from_response(self, response: Message) -> str: Returns: str: The objective extracted from the response, or empty string if not found. """ - if not response.message_pieces: - return "" - - piece = response.get_piece() - - if piece.api_role != "assistant": - return "" + from pyrit.score.message_scorer import extract_objective_from_previous_turn - conversation = self._memory.get_message_pieces(conversation_id=piece.conversation_id) - last_prompt = max(conversation, key=lambda x: x.sequence) - - # Every text message piece from the last turn - return "\n".join( - [ - piece.original_value - for piece in conversation - if piece.sequence == last_prompt.sequence - 1 and piece.original_value_data_type == "text" - ] + print_deprecation_message( + old_item="Scorer._extract_objective_from_response", + new_item="pyrit.score.message_scorer.extract_objective_from_previous_turn", + removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, ) + return extract_objective_from_previous_turn(message=response, memory=self._memory) @staticmethod async def score_response_async( @@ -850,20 +863,20 @@ async def score_response_async( skip_on_error_result=skip_on_error_result, ) obj_task = objective_scorer.score_async( - message=response, - objective=objective, - skip_on_error_result=skip_on_error_result, - role_filter=role_filter, + scorable=MessageScorable( + message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result + ), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) aux_scores, obj_scores = await asyncio.gather(aux_task, obj_task) result["auxiliary_scores"] = aux_scores result["objective_scores"] = obj_scores else: obj_scores = await objective_scorer.score_async( - message=response, - objective=objective, - skip_on_error_result=skip_on_error_result, - role_filter=role_filter, + scorable=MessageScorable( + message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result + ), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) result["objective_scores"] = obj_scores return result @@ -898,15 +911,9 @@ async def score_response_multiple_scorers_async( return [] # Create all scoring tasks, note TEMPORARY fix to prevent multi-piece responses from breaking scoring logic - tasks = [ - scorer.score_async( - message=response, - objective=objective, - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) - for scorer in scorers - ] + scorable = MessageScorable(message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result) + expectation = ScoringExpectation(objective=objective) if objective is not None else None + tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in scorers] if not tasks: return [] diff --git a/pyrit/score/scorer_evaluation/scorer_evaluator.py b/pyrit/score/scorer_evaluation/scorer_evaluator.py index 39902b3e59..b5672b9fbc 100644 --- a/pyrit/score/scorer_evaluation/scorer_evaluator.py +++ b/pyrit/score/scorer_evaluation/scorer_evaluator.py @@ -13,21 +13,15 @@ from scipy.stats import ttest_1samp from pyrit.common.path import SCORER_EVALS_PATH +from pyrit.score.message_scorer import extract_objective_from_previous_turn from pyrit.score.scorer_evaluation.human_labeled_dataset import ( HarmHumanLabeledEntry, HumanLabeledDataset, ObjectiveHumanLabeledEntry, ) from pyrit.score.scorer_evaluation.krippendorff import krippendorff_alpha -from pyrit.score.scorer_evaluation.metrics_type import ( - MetricsType, - RegistryUpdateBehavior, -) -from pyrit.score.scorer_evaluation.scorer_metrics import ( - HarmScorerMetrics, - ObjectiveScorerMetrics, - ScorerMetrics, -) +from pyrit.score.scorer_evaluation.metrics_type import MetricsType, RegistryUpdateBehavior +from pyrit.score.scorer_evaluation.scorer_metrics import HarmScorerMetrics, ObjectiveScorerMetrics, ScorerMetrics from pyrit.score.scorer_evaluation.scorer_metrics_io import ( find_harm_metrics_by_eval_hash, find_objective_metrics_by_eval_hash, @@ -381,6 +375,12 @@ async def evaluate_dataset_async( # Validate dataset and extract data assistant_responses, human_scores_list, objectives = self._validate_and_extract_data(labeled_dataset) + # Harm datasets carry no objective, so the previous turn stands in for one. + resolved_objectives = objectives or [ + extract_objective_from_previous_turn(message=response, memory=self.scorer._memory) + for response in assistant_responses + ] + # Transpose human scores so each row is a complete set of scores across all responses all_human_scores = np.array(human_scores_list).T @@ -392,9 +392,8 @@ async def evaluate_dataset_async( start_time = time.perf_counter() scores = await self.scorer.score_prompts_batch_async( messages=assistant_responses, - objectives=objectives, + objectives=resolved_objectives, batch_size=max_concurrency, - infer_objective_from_request=True, ) elapsed_time = time.perf_counter() - start_time total_scoring_time += elapsed_time @@ -534,7 +533,8 @@ def _validate_and_extract_data( Returns: Tuple of (assistant_responses, human_scores_list, None). - objectives is None for harm scoring (uses infer_objective_from_request). + objectives is None for harm scoring; the caller reads each objective from the + previous turn instead. Raises: ValueError: If dataset is not HARM type or has multiple harm categories. diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index 560f001b1f..d142809df1 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -7,11 +7,16 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score -from pyrit.score.float_scale.float_scale_score_aggregator import ( - FloatScaleAggregatorFunc, - FloatScaleScoreAggregator, +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, ) +from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleAggregatorFunc, FloatScaleScoreAggregator from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.score_utils import ORIGINAL_FLOAT_VALUE_KEY from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -102,9 +107,8 @@ async def _score_async( list[Score]: A list containing a single true/false Score object based on the threshold comparison. """ scores = await self._scorer.score_async( - message, - objective=objective, - role_filter=role_filter, + scorable=MessageScorable(message=message, role_filter=role_filter), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) # Aggregator handles 0-many scores and returns exactly one result (or raises if configured) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index 554749a92f..a8334ad520 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -7,7 +7,15 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, +) from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -98,7 +106,10 @@ async def _score_async( ValueError: If no scores are generated from the request response pieces. """ tasks = [ - scorer.score_async(message=message, objective=objective, role_filter=role_filter) + scorer.score_async( + scorable=MessageScorable(message=message, role_filter=role_filter), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, + ) for scorer in self._scorers ] diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index 5140d22afa..ef2af9072f 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -7,7 +7,15 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, +) from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -74,12 +82,9 @@ async def _score_async( list[Score]: A list containing a single Score object with the inverted true/false value. """ scores = await self._scorer.score_async( - message, - objective=objective, - role_filter=role_filter, + scorable=MessageScorable(message=message, role_filter=role_filter), + expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) - - # TrueFalseScorers only have a single score inv_score = scores[0] inv_score.score_value = str(True) if not inv_score.get_value() else str(False) diff --git a/pyrit/score/true_false/true_false_scorer.py b/pyrit/score/true_false/true_false_scorer.py index c3bb3399a7..36524dcbc9 100644 --- a/pyrit/score/true_false/true_false_scorer.py +++ b/pyrit/score/true_false/true_false_scorer.py @@ -6,11 +6,8 @@ from typing import TYPE_CHECKING from pyrit.models import Message, Score -from pyrit.score.scorer import Scorer -from pyrit.score.true_false.true_false_score_aggregator import ( - TrueFalseAggregatorFunc, - TrueFalseScoreAggregator, -) +from pyrit.score.message_scorer import MessageScorer +from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc, TrueFalseScoreAggregator if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget @@ -19,7 +16,7 @@ from pyrit.score.scorer_prompt_validator import ScorerPromptValidator -class TrueFalseScorer(Scorer): +class TrueFalseScorer(MessageScorer): """ Base class for scorers that return true/false binary scores. @@ -65,9 +62,7 @@ def __init__( # Set default evaluation file mapping if not already set by subclass if self.evaluation_file_mapping is None: - from pyrit.score.scorer_evaluation.scorer_evaluator import ( - ScorerEvalDatasetFiles, - ) + from pyrit.score.scorer_evaluation.scorer_evaluator import ScorerEvalDatasetFiles self.evaluation_file_mapping = ScorerEvalDatasetFiles( human_labeled_datasets_files=["objective/*.csv"], @@ -101,9 +96,7 @@ def get_scorer_metrics(self) -> ObjectiveScorerMetrics | None: ObjectiveScorerMetrics: The metrics for this scorer, or None if not found or not configured. """ from pyrit.common.path import SCORER_EVALS_PATH - from pyrit.score.scorer_evaluation.scorer_metrics_io import ( - find_objective_metrics_by_eval_hash, - ) + from pyrit.score.scorer_evaluation.scorer_metrics_io import find_objective_metrics_by_eval_hash if self.evaluation_file_mapping is None: return None diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo.py b/tests/unit/executor/attack/multi_turn/test_crescendo.py index e962dddf41..d0f3c865f0 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo.py @@ -1201,10 +1201,10 @@ async def test_check_refusal_does_not_skip_on_error_result( await attack._check_refusal_async(context=basic_context, objective="test task") - # Verify score_async was called with skip_on_error_result=False + # Verify the scorable does not skip error results mock_refusal_scorer.score_async.assert_called_once() - call_kwargs = mock_refusal_scorer.score_async.call_args.kwargs - assert call_kwargs.get("skip_on_error_result") is False, ( + scorable = mock_refusal_scorer.score_async.call_args.kwargs["scorable"] + assert scorable.skip_on_error_result is False, ( "Refusal scorer must be called with skip_on_error_result=False " "to ensure error responses are scored (treated as refusals) rather than skipped" ) diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py index 5352c0f939..2d52a943c9 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py @@ -19,14 +19,7 @@ CrescendoAttackContext, CrescendoAttackResult, ) -from pyrit.models import ( - AttackOutcome, - ComponentIdentifier, - ConversationType, - Message, - MessagePiece, - Score, -) +from pyrit.models import AttackOutcome, ComponentIdentifier, ConversationType, Message, MessagePiece, Score from pyrit.prompt_normalizer import PromptNormalizer from pyrit.score import Scorer, TrueFalseScorer @@ -286,7 +279,9 @@ async def score_objective(**_kwargs): assert result.conversation_id == final_conversation_id assert len({attempt.conversation_id for attempt in adversarial_target.attempts}) == 1 - refusal_inputs = [call.kwargs["message"].get_value() for call in refusal_scorer.score_async.await_args_list] + refusal_inputs = [ + call.kwargs["scorable"].message.get_value() for call in refusal_scorer.score_async.await_args_list + ] assert refusal_inputs == [ "response-1", "response-2", @@ -301,7 +296,7 @@ async def score_objective(**_kwargs): "response-9", "response-10-final", ] - assert [call.kwargs["objective"] for call in refusal_scorer.score_async.await_args_list] == [ + assert [call.kwargs["expectation"].objective for call in refusal_scorer.score_async.await_args_list] == [ f"question-{attempt}" for attempt in range(1, 13) ] objective_inputs = [call.kwargs["response"].get_value() for call in score_response.await_args_list] diff --git a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py index 239f7c444c..b23394d92c 100644 --- a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py +++ b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py @@ -37,6 +37,7 @@ ConversationType, Message, MessagePiece, + MessageScorable, Score, SeedPrompt, ) @@ -391,7 +392,7 @@ async def create_threshold_score_async(*, original_float_value: float, threshold ) # Score using the actual FloatScaleThresholdScorer - scores = await threshold_scorer.score_async(dummy_message) + scores = await threshold_scorer.score_async(scorable=MessageScorable(message=dummy_message)) return scores[0] @staticmethod diff --git a/tests/unit/models/test_scorable.py b/tests/unit/models/test_scorable.py new file mode 100644 index 0000000000..c8322ca4c7 --- /dev/null +++ b/tests/unit/models/test_scorable.py @@ -0,0 +1,144 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import dataclasses +import uuid + +import pytest + +from pyrit.models import ( + ContentScorable, + ConversationScorable, + MessagePiece, + MessageReferenceScorable, + MessageScorable, + OutputMatches, + ScoringExpectation, + ScoringScope, + SurfaceScorable, + ToolCalled, + ToolSequence, + TraceScorable, + Volatility, +) + + +def _message(): + return MessagePiece(role="assistant", original_value="hello").to_message() + + +@pytest.mark.parametrize( + "scorable, expected_volatility", + [ + (ConversationScorable(conversation_id="c1"), Volatility.APPEND_ONLY), + (MessageScorable(message=_message()), Volatility.STABLE), + (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), Volatility.STABLE), + (ContentScorable(value="hello"), Volatility.STABLE), + (SurfaceScorable(uri="/tmp/out.txt"), Volatility.MUTABLE), + (TraceScorable(trace_ids=("t1",)), Volatility.APPEND_ONLY), + ], +) +def test_scorable_volatility(scorable, expected_volatility): + assert scorable.volatility is expected_volatility + + +@pytest.mark.parametrize( + "scorable", + [ + ConversationScorable(conversation_id="c1"), + MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), + ContentScorable(value="hello"), + SurfaceScorable(uri="/tmp/out.txt"), + TraceScorable(trace_ids=("t1",)), + ], +) +def test_scorable_is_frozen(scorable): + with pytest.raises(dataclasses.FrozenInstanceError): + scorable.volatility = Volatility.MUTABLE + + +def test_conversation_scorable_defaults(): + scorable = ConversationScorable(conversation_id="c1") + + assert scorable.through_sequence is None + assert scorable.role_filter is None + + +def test_message_scorable_carries_the_message(): + message = _message() + scorable = MessageScorable(message=message, role_filter="assistant", skip_on_error_result=True) + + assert scorable.message is message + assert scorable.role_filter == "assistant" + assert scorable.skip_on_error_result is True + + +def test_message_reference_scorable_defaults(): + piece_id = uuid.uuid4() + scorable = MessageReferenceScorable(message_piece_ids=(piece_id,)) + + assert scorable.message_piece_ids == (piece_id,) + assert scorable.role_filter is None + assert scorable.skip_on_error_result is False + + +def test_content_scorable_defaults_to_text(): + assert ContentScorable(value="hello").data_type == "text" + + +def test_surface_scorable_defaults_to_file(): + scorable = SurfaceScorable(uri="/tmp/out.txt") + + assert scorable.surface == "file" + assert scorable.scope is None + + +def test_trace_scorable_accepts_a_scope(): + scope = ScoringScope(window="last_turn") + scorable = TraceScorable(trace_ids=("t1",), span_name="tool_call", scope=scope) + + assert scorable.scope is scope + assert scorable.span_name == "tool_call" + + +def test_scoring_scope_defaults(): + scope = ScoringScope() + + assert scope.window is None + assert scope.labels == {} + + +def test_expectation_defaults(): + expectation = ScoringExpectation() + + assert expectation.objective is None + assert expectation.conditions == () + assert expectation.extra == {} + + +def test_expectation_carries_conditions(): + conditions = (ToolCalled(name="send_email"), ToolSequence(names=("a", "b")), OutputMatches(value="42")) + expectation = ScoringExpectation(objective="exfiltrate", conditions=conditions) + + assert expectation.objective == "exfiltrate" + assert expectation.conditions == conditions + + +def test_expectation_is_frozen(): + expectation = ScoringExpectation(objective="exfiltrate") + + with pytest.raises(dataclasses.FrozenInstanceError): + expectation.objective = "something else" + + +def test_tool_called_defaults(): + condition = ToolCalled(name="send_email") + + assert condition.name == "send_email" + assert condition.arguments is None + + +def test_expectations_with_equal_values_compare_equal(): + assert ScoringExpectation(objective="a", conditions=(ToolCalled(name="t"),)) == ScoringExpectation( + objective="a", conditions=(ToolCalled(name="t"),) + ) diff --git a/tests/unit/score/test_azure_content_filter.py b/tests/unit/score/test_azure_content_filter.py index 16e759de6e..da59ed00f2 100644 --- a/tests/unit/score/test_azure_content_filter.py +++ b/tests/unit/score/test_azure_content_filter.py @@ -8,15 +8,11 @@ import pytest from azure.ai.contentsafety.models import TextCategory -from unit.mocks import ( - get_audio_message_piece, - get_image_message_piece, - get_test_message_piece, -) +from unit.mocks import get_audio_message_piece, get_image_message_piece, get_test_message_piece from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece +from pyrit.models import Message, MessagePiece, MessageScorable from pyrit.score.float_scale.azure_content_filter_scorer import AzureContentFilterScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer @@ -46,7 +42,7 @@ async def test_score_async_unsupported_data_type_returns_zero( # Unified FloatScaleScorer fallback: when all pieces are filtered out, return a single # Score(0.0) instead of an empty list (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(message=request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 @@ -338,7 +334,7 @@ async def test_azure_content_filter_scorer_blocked_returns_one_score_per_categor ) message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(message=message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 2 assert {s.score_category[0] for s in scores} == {TextCategory.HATE.value, TextCategory.VIOLENCE.value} @@ -363,7 +359,7 @@ async def test_azure_content_filter_scorer_blocked_default_categories_returns_fo ) message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(message=message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 4 assert {s.score_category[0] for s in scores} == {c.value for c in TextCategory} diff --git a/tests/unit/score/test_conversation_history_scorer.py b/tests/unit/score/test_conversation_history_scorer.py index e60e622948..fb4be72386 100644 --- a/tests/unit/score/test_conversation_history_scorer.py +++ b/tests/unit/score/test_conversation_history_scorer.py @@ -7,12 +7,8 @@ import pytest from pyrit.memory import CentralMemory -from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score -from pyrit.score import ( - Scorer, - SelfAskGeneralFloatScaleScorer, - create_conversation_scorer, -) +from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score +from pyrit.score import Scorer, SelfAskGeneralFloatScaleScorer, create_conversation_scorer from pyrit.score.conversation_scorer import ConversationScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -143,7 +139,7 @@ async def test_conversation_history_scorer_score_async_success(patch_central_dat mock_scorer.validate_return_scores = MagicMock() scorer = create_conversation_scorer(scorer=mock_scorer) - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 1 result_score = scores[0] @@ -183,7 +179,7 @@ async def test_conversation_history_scorer_conversation_not_found(patch_central_ message.message_pieces = [message_piece] with pytest.raises(RuntimeError, match=f"Conversation with ID {nonexistent_conversation_id} not found in memory"): - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) async def test_conversation_history_scorer_filters_roles_correctly(patch_central_database): @@ -233,7 +229,7 @@ async def test_conversation_history_scorer_filters_roles_correctly(patch_central mock_scorer.validate_return_scores = MagicMock() scorer = create_conversation_scorer(scorer=mock_scorer) - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) call_args = mock_scorer._score_async.call_args called_message = call_args.kwargs["message"] @@ -278,7 +274,7 @@ async def test_conversation_history_scorer_preserves_metadata(patch_central_data scorer = create_conversation_scorer(scorer=mock_scorer) - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) call_args = mock_scorer._score_async.call_args called_message = call_args.kwargs["message"] @@ -329,7 +325,7 @@ async def test_conversation_scorer_persists_scores_exactly_once(patch_central_da conv_scorer = create_conversation_scorer(scorer=mock_scorer) message = MagicMock() message.message_pieces = [message_piece] - result_scores = await conv_scorer.score_async(message) + result_scores = await conv_scorer.score_async(scorable=MessageScorable(message=message)) assert len(result_scores) == 1 assert result_scores[0].id == original_id, ( @@ -549,7 +545,7 @@ async def test_conversation_scorer_uses_partial_content_when_score_blocked_conte scorer = create_conversation_scorer(scorer=mock_scorer) scorer.score_blocked_content = True - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 1 @@ -621,7 +617,7 @@ async def test_conversation_scorer_uses_error_json_when_score_blocked_content_di scorer = create_conversation_scorer(scorer=mock_scorer) # score_blocked_content defaults to False - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 1 @@ -689,7 +685,7 @@ async def test_conversation_scorer_blocked_input_message_does_not_raise(patch_ce scorer = create_conversation_scorer(scorer=mock_scorer) # Must not raise — previously raised ValueError on the blocked piece. - scores = await scorer.score_async(blocked_message) + scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) assert len(scores) == 1 mock_scorer._score_async.assert_awaited_once() @@ -783,7 +779,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st inner_scorer = HarmfulContentDetector() scorer = create_conversation_scorer(scorer=inner_scorer) - scores = await scorer.score_async(blocked_message) + scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) assert len(scores) == 1 # Must be 1.0 (real score from prior turns), NOT 0.0 (fallback from rejected synthetic piece) diff --git a/tests/unit/score/test_float_scale_threshold_scorer.py b/tests/unit/score/test_float_scale_threshold_scorer.py index b98cb183d8..7e47e30c14 100644 --- a/tests/unit/score/test_float_scale_threshold_scorer.py +++ b/tests/unit/score/test_float_scale_threshold_scorer.py @@ -7,7 +7,7 @@ import pytest from pyrit.memory import CentralMemory, MemoryInterface -from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score +from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score from pyrit.score import FloatScaleThresholdScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -170,9 +170,7 @@ async def test_float_scale_threshold_scorer_with_raise_on_empty_aggregator(): Test that FloatScaleThresholdScorer raises ValueError when using RAISE_ON_EMPTY aggregator and the underlying scorer returns no scores. """ - from pyrit.score.float_scale.float_scale_score_aggregator import ( - FloatScaleScoreAggregator, - ) + from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleScoreAggregator memory = MagicMock(MemoryInterface) @@ -262,7 +260,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st ) blocked_message = Message(message_pieces=[blocked_piece]) - scores = await threshold_scorer.score_async(blocked_message) + scores = await threshold_scorer.score_async(scorable=MessageScorable(message=blocked_message)) assert len(scores) == 1 binary_score = scores[0] diff --git a/tests/unit/score/test_gandalf_scorer.py b/tests/unit/score/test_gandalf_scorer.py index e47a4c39cd..45efd3da7a 100644 --- a/tests/unit/score/test_gandalf_scorer.py +++ b/tests/unit/score/test_gandalf_scorer.py @@ -9,7 +9,7 @@ from pyrit.exceptions.exception_classes import PyritException from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece +from pyrit.models import Message, MessagePiece, MessageScorable from pyrit.prompt_target import GandalfLevel from pyrit.score import GandalfScorer @@ -64,7 +64,7 @@ async def test_gandalf_scorer_score( mocked_post.return_value = MagicMock(json=lambda: {"success": password_correct, "message": "Message"}) - scores = await scorer.score_async(response) + scores = await scorer.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].get_value() == password_correct @@ -99,7 +99,7 @@ async def test_gandalf_scorer_set_system_prompt( mocked_post.return_value = MagicMock(json=lambda: {"success": True, "message": "Message"}) - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) chat_target.set_system_prompt.assert_called_once() @@ -124,7 +124,7 @@ async def test_gandalf_scorer_adds_to_memory(mocked_post, level: GandalfLevel, s with patch.object(sqlite_instance, "get_message_pieces", return_value=[generated_request.message_pieces[0]]): scorer = GandalfScorer(level=level, chat_target=chat_target) - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) @pytest.mark.parametrize("level", [GandalfLevel.LEVEL_1, GandalfLevel.LEVEL_2, GandalfLevel.LEVEL_3]) @@ -140,7 +140,7 @@ async def test_gandalf_scorer_runtime_error_retries(level: GandalfLevel, sqlite_ scorer = GandalfScorer(level=level, chat_target=chat_target) with pytest.raises(PyritException, match="Error in scorer GandalfScorer"): - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) assert chat_target.send_prompt_async.call_count == 1 @@ -167,4 +167,4 @@ async def test_gandalf_scorer_wraps_httpx_error_as_pyrit_exception(mocked_post, scorer = GandalfScorer(level=GandalfLevel.LEVEL_1, chat_target=chat_target) with pytest.raises(PyritException, match="Error in scorer GandalfScorer"): - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) diff --git a/tests/unit/score/test_insecure_code_scorer.py b/tests/unit/score/test_insecure_code_scorer.py index e0370a196d..8404e7f746 100644 --- a/tests/unit/score/test_insecure_code_scorer.py +++ b/tests/unit/score/test_insecure_code_scorer.py @@ -6,7 +6,15 @@ import pytest from pyrit.exceptions.exception_classes import InvalidJsonException -from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, SeedPrompt, UnvalidatedScore +from pyrit.models import ( + ComponentIdentifier, + Message, + MessagePiece, + MessageScorable, + Score, + SeedPrompt, + UnvalidatedScore, +) from pyrit.prompt_target import PromptTarget from pyrit.score import InsecureCodeScorer @@ -47,7 +55,7 @@ async def test_insecure_code_scorer_valid_response(mock_chat_target): message = MessagePiece(role="user", original_value="sample code").to_message() # Call the score_async method - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) # Assertions assert len(scores) == 1 @@ -70,7 +78,7 @@ async def test_insecure_code_scorer_invalid_json(mock_chat_target): message = MessagePiece(role="user", original_value="sample code").to_message() with pytest.raises(InvalidJsonException, match="Error in scorer InsecureCodeScorer.*Invalid JSON"): - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) # Ensure memory functions were not called mock_add_scores.assert_not_called() @@ -109,7 +117,7 @@ async def test_score_async_unsupported_data_type_returns_zero(mock_chat_target, # Unified FloatScaleScorer fallback: returns a single Score(0.0) when all pieces are filtered # out (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py new file mode 100644 index 0000000000..0655a84c33 --- /dev/null +++ b/tests/unit/score/test_message_scorer.py @@ -0,0 +1,377 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import uuid + +import pytest + +from pyrit.memory import MemoryInterface +from pyrit.models import ( + ComponentIdentifier, + ContentScorable, + ConversationScorable, + Message, + MessagePiece, + MessageReferenceScorable, + MessageScorable, + Score, + ScoringExpectation, +) +from pyrit.score import ScorerPromptValidator, TrueFalseScorer +from pyrit.score.message_scorer import extract_objective_from_previous_turn + + +class PermissiveValidator(ScorerPromptValidator): + def validate(self, message, objective=None): + pass + + def is_message_piece_supported(self, message_piece): + return True + + +class RecordingScorer(TrueFalseScorer): + """A message scorer that remembers what it was asked to score.""" + + def __init__(self): + super().__init__(validator=PermissiveValidator()) + self.scored_messages: list[Message] = [] + self.scored_objectives: list[str | None] = [] + + def _build_identifier(self) -> ComponentIdentifier: + return self._create_identifier() + + async def _score_async(self, message: Message, *, objective: str | None = None) -> list[Score]: + self.scored_messages.append(message) + self.scored_objectives.append(objective) + return [self._build_score(message.get_piece(), objective)] + + async def _score_piece_async(self, message_piece: MessagePiece, *, objective: str | None = None) -> list[Score]: + return [self._build_score(message_piece, objective)] + + def _build_score(self, message_piece: MessagePiece, objective: str | None) -> Score: + return Score( + score_value="true", + score_value_description="desc", + score_type="true_false", + score_category=None, + score_metadata=None, + score_rationale="rationale", + scorer_class_identifier=self.get_identifier(), + message_piece_id=message_piece.id, + objective=objective, + ) + + +def _assistant_message(value: str = "response", conversation_id: str | None = None) -> Message: + return MessagePiece( + role="assistant", + original_value=value, + conversation_id=conversation_id or str(uuid.uuid4()), + ).to_message() + + +@pytest.mark.usefixtures("patch_central_database") +class TestScorableResolution: + """MessageScorer reduces every message-shaped scorable to a single Message.""" + + async def test_message_scorable_is_scored_directly(self): + scorer = RecordingScorer() + message = _assistant_message() + + scores = await scorer.score_async(scorable=MessageScorable(message=message)) + + assert len(scores) == 1 + assert scorer.scored_messages == [message] + + async def test_message_reference_scorable_resolves_from_memory(self, sqlite_instance: MemoryInterface): + message = _assistant_message("stored response") + sqlite_instance.add_message_to_memory(request=message) + piece_id = message.get_piece().id + scorer = RecordingScorer() + + scores = await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(piece_id,))) + + assert len(scores) == 1 + assert scorer.scored_messages[0].get_value() == "stored response" + + async def test_message_reference_scorable_not_in_memory_raises(self): + scorer = RecordingScorer() + missing_id = uuid.uuid4() + + with pytest.raises(ValueError, match="No message pieces found in memory"): + await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(missing_id,))) + + async def test_message_reference_scorable_spanning_messages_raises(self, sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + first = MessagePiece( + role="user", original_value="ask", conversation_id=conversation_id, sequence=0 + ).to_message() + second = MessagePiece( + role="assistant", original_value="answer", conversation_id=conversation_id, sequence=1 + ).to_message() + sqlite_instance.add_message_to_memory(request=first) + sqlite_instance.add_message_to_memory(request=second) + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="exactly one message"): + await scorer.score_async( + scorable=MessageReferenceScorable( + message_piece_ids=(first.get_piece().id, second.get_piece().id), + ) + ) + + async def test_content_scorable_is_never_persisted(self): + scorer = RecordingScorer() + + scores = await scorer.score_async(scorable=ContentScorable(value="loose text")) + + scored_piece = scorer.scored_messages[0].get_piece() + assert scored_piece.original_value == "loose text" + assert scored_piece.role == "user" + assert scored_piece.not_in_memory is True + # Memory cannot link a score to a piece it never stored. + assert scores[0].message_piece_id is None + + async def test_unsupported_scorable_raises_type_error(self): + scorer = RecordingScorer() + + with pytest.raises(TypeError, match="cannot score ConversationScorable"): + await scorer.score_async(scorable=ConversationScorable(conversation_id=str(uuid.uuid4()))) + + +@pytest.mark.usefixtures("patch_central_database") +class TestScorableFilters: + """role_filter and skip_on_error_result are fields on the scorable, not call parameters.""" + + async def test_role_filter_mismatch_skips_scoring(self): + scorer = RecordingScorer() + message = _assistant_message() + + scores = await scorer.score_async(scorable=MessageScorable(message=message, role_filter="user")) + + assert scores == [] + assert scorer.scored_messages == [] + + async def test_role_filter_match_scores(self): + scorer = RecordingScorer() + message = _assistant_message() + + scores = await scorer.score_async(scorable=MessageScorable(message=message, role_filter="assistant")) + + assert len(scores) == 1 + + async def test_skip_on_error_result_skips_error_message(self): + scorer = RecordingScorer() + message = MessagePiece( + role="assistant", + original_value="blocked", + original_value_data_type="error", + response_error="blocked", + ).to_message() + + scores = await scorer.score_async(scorable=MessageScorable(message=message, skip_on_error_result=True)) + + assert scores == [] + assert scorer.scored_messages == [] + + async def test_error_message_is_scored_when_not_skipping(self): + scorer = RecordingScorer() + message = MessagePiece( + role="assistant", + original_value="blocked", + original_value_data_type="error", + response_error="blocked", + ).to_message() + + scores = await scorer.score_async(scorable=MessageScorable(message=message)) + + assert len(scores) == 1 + + +@pytest.mark.usefixtures("patch_central_database") +class TestExpectation: + """The expectation carries what to look for.""" + + async def test_objective_reaches_the_scorer(self): + scorer = RecordingScorer() + + await scorer.score_async( + scorable=MessageScorable(message=_assistant_message()), + expectation=ScoringExpectation(objective="find the objective"), + ) + + assert scorer.scored_objectives == ["find the objective"] + + async def test_no_expectation_means_no_objective(self): + scorer = RecordingScorer() + + await scorer.score_async(scorable=MessageScorable(message=_assistant_message())) + + assert scorer.scored_objectives == [None] + + +@pytest.mark.usefixtures("patch_central_database") +class TestDeprecatedParameters: + """The legacy message-shaped parameters survive one release behind a warning.""" + + async def test_positional_message_maps_to_message_scorable(self): + scorer = RecordingScorer() + message = _assistant_message() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + scores = await scorer.score_async(message) + + assert len(scores) == 1 + assert scorer.scored_messages == [message] + + async def test_keyword_message_maps_to_message_scorable(self): + scorer = RecordingScorer() + message = _assistant_message() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async(message=message) + + assert scorer.scored_messages == [message] + + async def test_message_does_not_widen_to_the_stored_conversation(self, sqlite_instance: MemoryInterface): + """A supplied message must map to MessageScorable, never ConversationScorable.""" + conversation_id = str(uuid.uuid4()) + sqlite_instance.add_message_to_memory( + request=MessagePiece( + role="user", + original_value="an earlier turn that must not be scored", + conversation_id=conversation_id, + sequence=0, + ).to_message() + ) + message = _assistant_message("only this turn", conversation_id=conversation_id) + sqlite_instance.add_message_to_memory(request=message) + scorer = RecordingScorer() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async(message) + + assert [scored.get_value() for scored in scorer.scored_messages] == ["only this turn"] + + async def test_objective_maps_to_expectation(self): + scorer = RecordingScorer() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async(_assistant_message(), objective="legacy objective") + + assert scorer.scored_objectives == ["legacy objective"] + + async def test_legacy_role_filter_maps_onto_the_scorable(self): + scorer = RecordingScorer() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + scores = await scorer.score_async(_assistant_message(), role_filter="user") + + assert scores == [] + + async def test_legacy_skip_on_error_result_maps_onto_the_scorable(self): + scorer = RecordingScorer() + message = MessagePiece( + role="assistant", + original_value="blocked", + original_value_data_type="error", + response_error="blocked", + ).to_message() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + scores = await scorer.score_async(message, skip_on_error_result=True) + + assert scores == [] + + async def test_infer_objective_from_request_reads_the_previous_turn(self, sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + sqlite_instance.add_message_to_memory( + request=MessagePiece( + role="user", + original_value="the inferred objective", + conversation_id=conversation_id, + sequence=0, + ).to_message() + ) + message = _assistant_message("response", conversation_id=conversation_id) + sqlite_instance.add_message_to_memory(request=message) + scorer = RecordingScorer() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async(message, infer_objective_from_request=True) + + assert scorer.scored_objectives == ["the inferred objective"] + + async def test_new_signature_emits_no_warning(self, recwarn): + scorer = RecordingScorer() + + await scorer.score_async( + scorable=MessageScorable(message=_assistant_message()), + expectation=ScoringExpectation(objective="objective"), + ) + + assert [warning for warning in recwarn if issubclass(warning.category, DeprecationWarning)] == [] + + +@pytest.mark.usefixtures("patch_central_database") +class TestConflictingInputs: + """The shim refuses input it cannot map without guessing.""" + + async def test_message_and_scorable_together_raises(self): + scorer = RecordingScorer() + message = _assistant_message() + + with pytest.raises(ValueError, match="not both"): + await scorer.score_async(message, scorable=MessageScorable(message=message)) + + async def test_neither_message_nor_scorable_raises(self): + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="must be provided"): + await scorer.score_async() + + async def test_objective_and_expectation_together_raises(self): + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="not both"): + await scorer.score_async( + scorable=MessageScorable(message=_assistant_message()), + objective="one", + expectation=ScoringExpectation(objective="two"), + ) + + @pytest.mark.parametrize("kwargs", [{"role_filter": "assistant"}, {"skip_on_error_result": True}]) + async def test_message_flags_with_a_scorable_raises(self, kwargs): + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="fields on the message scorable"): + await scorer.score_async(scorable=MessageScorable(message=_assistant_message()), **kwargs) + + +@pytest.mark.usefixtures("patch_central_database") +class TestExtractObjectiveFromPreviousTurn: + """The objective lookup belongs to whoever builds the expectation.""" + + def test_reads_the_turn_before_the_response(self, sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + sqlite_instance.add_message_to_memory( + request=MessagePiece( + role="user", original_value="the request", conversation_id=conversation_id, sequence=0 + ).to_message() + ) + message = _assistant_message("the response", conversation_id=conversation_id) + sqlite_instance.add_message_to_memory(request=message) + + objective = extract_objective_from_previous_turn(message=message, memory=sqlite_instance) + + assert objective == "the request" + + def test_returns_empty_for_a_user_message(self, sqlite_instance: MemoryInterface): + message = MessagePiece(role="user", original_value="a request").to_message() + + assert extract_objective_from_previous_turn(message=message, memory=sqlite_instance) == "" + + def test_returns_empty_when_the_conversation_is_not_stored(self, sqlite_instance: MemoryInterface): + message = _assistant_message() + + assert extract_objective_from_previous_turn(message=message, memory=sqlite_instance) == "" diff --git a/tests/unit/score/test_plagiarism_scorer.py b/tests/unit/score/test_plagiarism_scorer.py index aef799eea6..ecf612ec8d 100644 --- a/tests/unit/score/test_plagiarism_scorer.py +++ b/tests/unit/score/test_plagiarism_scorer.py @@ -7,11 +7,8 @@ from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece -from pyrit.score import ( - PlagiarismMetric, - PlagiarismScorer, -) +from pyrit.models import MessagePiece, MessageScorable +from pyrit.score import PlagiarismMetric, PlagiarismScorer @pytest.mark.usefixtures("patch_central_database") @@ -55,7 +52,7 @@ async def test_score_async_lcs_metric(self): request = message_piece.to_message() - scores = await scorer.score_async(message=request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 score = scores[0] @@ -185,7 +182,7 @@ async def test_score_async_adds_to_memory(self): ).to_message() with patch.object(CentralMemory, "get_memory_instance", return_value=memory): - await scorer.score_async(request) + await scorer.score_async(scorable=MessageScorable(message=request)) memory.add_scores_to_memory.assert_called_once() async def test_score_async_unsupported_data_type_returns_zero(self, patch_central_database): @@ -202,7 +199,7 @@ async def test_score_async_unsupported_data_type_returns_zero(self, patch_centra # Unified FloatScaleScorer fallback: returns a single Score(0.0) when all pieces are filtered # out (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 diff --git a/tests/unit/score/test_question_answer_scorer.py b/tests/unit/score/test_question_answer_scorer.py index 4c3519d6b0..5113f73da3 100644 --- a/tests/unit/score/test_question_answer_scorer.py +++ b/tests/unit/score/test_question_answer_scorer.py @@ -10,7 +10,7 @@ from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece +from pyrit.models import Message, MessagePiece, MessageScorable from pyrit.score import QuestionAnswerScorer @@ -39,7 +39,7 @@ async def test_score_async_unsupported_image_type_returns_false( message = Message(message_pieces=[image_message_piece]) # With raise_on_no_valid_pieces=False (default), returns False for unsupported data types - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale @@ -58,7 +58,7 @@ async def test_score_async_missing_metadata_returns_false(patch_central_database scorer = QuestionAnswerScorer(category=["new_category"]) # With raise_on_no_valid_pieces=False (default), returns False for missing metadata - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale @@ -80,7 +80,7 @@ async def test_question_answer_scorer_score(response: str, expected_score: bool, scorer = QuestionAnswerScorer(category=["new_category"]) message = Message(message_pieces=[text_message_piece]) - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert len(scores) == 1 result_score = scores[0] @@ -101,7 +101,7 @@ async def test_question_answer_scorer_adds_to_memory(): prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, ).to_message() - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) memory.add_scores_to_memory.assert_called_once() @@ -117,6 +117,6 @@ async def test_question_answer_scorer_no_category(): converted_value_data_type="text", prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, ).to_message() - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) memory.add_scores_to_memory.assert_called_once() diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index f1bcd4d000..ad9c4e70f7 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -10,17 +10,19 @@ from unit.mocks import get_mock_target_identifier from pyrit.exceptions import InvalidJsonException, remove_markdown_json -from pyrit.memory import CentralMemory -from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score +from pyrit.memory import MemoryInterface +from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score, ScoringExpectation from pyrit.prompt_target import PromptTarget from pyrit.score import ( FloatScaleScorer, JsonSchemaResponseHandler, + MessageScorer, Scorer, ScorerPromptValidator, TrueFalseScorer, ) from pyrit.score.llm_scoring import _run_llm_scoring_async +from pyrit.score.message_scorer import extract_objective_from_previous_turn @pytest.fixture @@ -113,7 +115,7 @@ def __init__(self, *, enforce_all_pieces_valid: bool = False, raise_on_no_valid_ ) -class MockFloatScorer(Scorer): +class MockFloatScorer(MessageScorer): """Mock scorer that tracks which pieces were scored.""" def __init__(self, *, validator: ScorerPromptValidator): @@ -406,13 +408,12 @@ async def test_score_value_with_llm_prepended_text_works_with_audio(good_json, p assert audio_piece.original_value == str(audio_path) -def test_scorer_extract_task_from_response(patch_central_database): +def test_extract_objective_from_previous_turn(patch_central_database): """ - Test that _extract_task_from_response properly gathers text from the + Test that extract_objective_from_previous_turn properly gathers text from the last turn. We'll mock out the memory's get_message_pieces method. """ - scorer = MockScorer() - mock_memory = MagicMock() + mock_memory = MagicMock(spec=MemoryInterface) response_piece = MessagePiece(original_value="og prompt", role="assistant", conversation_id="xyz", sequence=2) @@ -428,9 +429,8 @@ def test_scorer_extract_task_from_response(patch_central_database): response_piece, ] - with patch.object(CentralMemory, "get_memory_instance", return_value=mock_memory): - extracted_task = scorer._extract_objective_from_response(response_piece.to_message()) - assert "User's question about the universe" in extracted_task + extracted_task = extract_objective_from_previous_turn(message=response_piece.to_message(), memory=mock_memory) + assert "User's question about the universe" in extracted_task async def test_scorer_score_responses_batch_async(patch_central_database): @@ -447,9 +447,7 @@ async def test_scorer_score_responses_batch_async(patch_central_database): user_req = MessagePiece(role="user", original_value="Hello user", sequence=1).to_message() assistant_resp = MessagePiece(role="assistant", original_value="Hello from assistant", sequence=2).to_message() - results = await scorer.score_prompts_batch_async( - messages=[user_req, assistant_resp], batch_size=10, infer_objective_from_request=True - ) + results = await scorer.score_prompts_batch_async(messages=[user_req, assistant_resp], batch_size=10) # Verify mock_score_async was called twice assert mock_score_async.call_count == 2 @@ -457,10 +455,8 @@ async def test_scorer_score_responses_batch_async(patch_central_database): # Get the call_args for the first call _, first_call_kwargs = mock_score_async.call_args_list[0] - assert "message" in first_call_kwargs - assert "objective" in first_call_kwargs - assert "infer_objective_from_request" in first_call_kwargs - assert first_call_kwargs["message"] == user_req + assert first_call_kwargs["scorable"] == MessageScorable(message=user_req) + assert first_call_kwargs["expectation"] == ScoringExpectation(objective="") assert fake_scores[0] in results assert len(fake_scores) == 2 @@ -494,7 +490,7 @@ async def test_score_prompts_batch_async_defaults_objectives_when_none(patch_cen await scorer.score_prompts_batch_async(messages=[message]) _, call_kwargs = mock_score_async.call_args - assert call_kwargs["objective"] == "" + assert call_kwargs["expectation"] == ScoringExpectation(objective="") async def test_score_image_batch_async_works_when_objectives_none(patch_central_database): @@ -571,17 +567,14 @@ async def test_score_response_async_parallel_execution(): assert score1_1 in result["auxiliary_scores"] assert score2_1 in result["auxiliary_scores"] + expected_scorable = MessageScorable(message=response, role_filter="assistant", skip_on_error_result=True) scorer1.score_async.assert_any_call( - message=response, - objective="test task", - role_filter="assistant", - skip_on_error_result=True, + scorable=expected_scorable, + expectation=ScoringExpectation(objective="test task"), ) scorer2.score_async.assert_any_call( - message=response, - objective="test task", - role_filter="assistant", - skip_on_error_result=True, + scorable=expected_scorable, + expectation=ScoringExpectation(objective="test task"), ) @@ -600,7 +593,10 @@ async def test_score_async_no_matching_role(): """Test that score_response_select_first_success_async returns None when no pieces match role filter.""" response = Message(message_pieces=[MessagePiece(role="user", original_value="test", conversation_id="test-convo")]) scorer = MockScorer() - result = await scorer.score_async(message=response, role_filter="assistant", objective="test task") + result = await scorer.score_async( + scorable=MessageScorable(message=response, role_filter="assistant"), + expectation=ScoringExpectation(objective="test task"), + ) assert result == [] @@ -691,14 +687,14 @@ async def test_score_response_success_async_parallel_scoring_per_piece(): # Track call order call_order = [] - async def mock_score_async_1(message: Message, **kwargs) -> list[Score]: - call_order.append(("scorer1", message.message_pieces[0].original_value)) + async def mock_score_async_1(*, scorable: MessageScorable, **kwargs) -> list[Score]: + call_order.append(("scorer1", scorable.message.message_pieces[0].original_value)) score = MagicMock(spec=Score) score.get_value.return_value = False return [score] - async def mock_score_async_2(message: Message, **kwargs) -> list[Score]: - call_order.append(("scorer2", message.message_pieces[0].original_value)) + async def mock_score_async_2(*, scorable: MessageScorable, **kwargs) -> list[Score]: + call_order.append(("scorer2", scorable.message.message_pieces[0].original_value)) score = MagicMock(spec=Score) score.get_value.return_value = False return [score] @@ -966,14 +962,14 @@ async def test_score_response_async_concurrent_execution(): # Track call order to verify concurrent execution call_order = [] - async def mock_aux_score_async(message: Message, **kwargs) -> list[Score]: + async def mock_aux_score_async(**kwargs) -> list[Score]: call_order.append("aux_start") # Yield so the other scorer can interleave (proves concurrent execution). await asyncio.sleep(0) call_order.append("aux_end") return [MagicMock(spec=Score)] - async def mock_obj_score_async(message: Message, **kwargs) -> list[Score]: + async def mock_obj_score_async(**kwargs) -> list[Score]: call_order.append("obj_start") # Yield so the other scorer can interleave (proves concurrent execution). await asyncio.sleep(0) @@ -1053,7 +1049,7 @@ async def test_get_supported_pieces_filters_unsupported_data_types(patch_central response = Message(message_pieces=[text_piece, image_piece, audio_piece]) # Score the response - scores = await scorer.score_async(response) + scores = await scorer.score_async(scorable=MessageScorable(message=response)) # Should only score the text piece assert len(scorer.scored_piece_ids) == 1 @@ -1087,7 +1083,7 @@ async def test_unsupported_pieces_ignored_when_enforce_all_pieces_valid_false(pa response = Message(message_pieces=[image_piece, text_piece]) # Should not raise an error, just skip the image piece - scores = await scorer.score_async(response) + scores = await scorer.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert len(scorer.scored_piece_ids) == 1 @@ -1119,7 +1115,7 @@ async def test_all_unsupported_pieces_raises_error(patch_central_database): # Should raise error from validator because no valid pieces to score with pytest.raises(ValueError, match="There are no valid pieces to score"): - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) # No pieces should have been scored assert len(scorer.scored_piece_ids) == 0 @@ -1176,7 +1172,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st response = Message(message_pieces=[text_piece, image_piece]) # Score the response - scores = await scorer.score_async(response) + scores = await scorer.score_async(scorable=MessageScorable(message=response)) # Should only score the text piece assert len(scorer.scored_piece_ids) == 1 @@ -1212,7 +1208,7 @@ async def test_base_scorer_score_async_implementation(patch_central_database): response = Message(message_pieces=[text_piece1, text_piece2]) # Score the response - scores = await scorer.score_async(response) + scores = await scorer.score_async(scorable=MessageScorable(message=response)) # Should score both pieces assert len(scorer.scored_piece_ids) == 2 @@ -1344,7 +1340,7 @@ async def test_blocked_response_returns_specific_rationale( ) response = Message(message_pieces=[blocked_piece]) - scores = await true_false_scorer_returns_empty.score_async(response) + scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1366,7 +1362,7 @@ async def test_error_response_returns_specific_rationale( ) response = Message(message_pieces=[error_piece]) - scores = await true_false_scorer_returns_empty.score_async(response) + scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1388,7 +1384,7 @@ async def test_filtered_pieces_returns_generic_rationale( ) response = Message(message_pieces=[normal_piece]) - scores = await true_false_scorer_returns_empty.score_async(response) + scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1411,7 +1407,7 @@ async def test_blocked_takes_precedence_over_generic_error( ) response = Message(message_pieces=[blocked_piece]) - scores = await true_false_scorer_returns_empty.score_async(response) + scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) # Should specifically mention blocked, not generic error assert "blocked" in scores[0].score_rationale.lower() @@ -1470,7 +1466,7 @@ async def test_blocked_response_returns_zero_with_blocked_rationale( ) response = Message(message_pieces=[blocked_piece]) - scores = await float_scale_scorer_returns_empty.score_async(response) + scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].score_type == "float_scale" @@ -1492,7 +1488,7 @@ async def test_other_error_response_returns_zero_with_error_rationale( ) response = Message(message_pieces=[error_piece]) - scores = await float_scale_scorer_returns_empty.score_async(response) + scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].get_value() == 0.0 @@ -1513,7 +1509,7 @@ async def test_filtered_pieces_return_zero_with_generic_rationale( ) response = Message(message_pieces=[normal_piece]) - scores = await float_scale_scorer_returns_empty.score_async(response) + scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) assert len(scores) == 1 assert scores[0].get_value() == 0.0 @@ -1539,7 +1535,7 @@ async def test_text_only_scorer_filters_blocked_via_validator( with patch.object( float_scale_scorer_returns_empty, "_score_piece_async", new_callable=AsyncMock ) as mock_score_piece: - scores = await float_scale_scorer_returns_empty.score_async(response) + scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) mock_score_piece.assert_not_called() assert len(scores) == 1 @@ -1821,14 +1817,14 @@ async def test_raises_by_default(self): scorer = _ForwarderTrueFalseScorer(chat_target=_make_scorer_blocking_target()) with pytest.raises(ScorerLLMResponseBlockedException, match="blocked by content filtering"): - await scorer.score_async(_make_normal_input_message()) + await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) async def test_returns_false_when_flag_disabled(self): target = _make_scorer_blocking_target() scorer = _ForwarderTrueFalseScorer(chat_target=target) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(_make_normal_input_message()) + scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1840,7 +1836,7 @@ async def test_returns_zero_for_float_scale_when_flag_disabled(self): scorer = _ForwarderFloatScaleScorer(chat_target=_make_scorer_blocking_target()) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(_make_normal_input_message()) + scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) assert len(scores) == 1 assert scores[0].score_value == "0.0" @@ -1852,13 +1848,13 @@ async def test_direct_transport_caller_raises_by_default(self): scorer = _DirectTransportTrueFalseScorer(chat_target=_make_scorer_blocking_target()) with pytest.raises(ScorerLLMResponseBlockedException, match="blocked by content filtering"): - await scorer.score_async(_make_normal_input_message()) + await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) async def test_direct_transport_caller_returns_false_when_flag_disabled(self): scorer = _DirectTransportTrueFalseScorer(chat_target=_make_scorer_blocking_target()) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(_make_normal_input_message()) + scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2062,7 +2058,7 @@ async def test_default_false_skips_blocked_piece_text_only_scorer(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2074,7 +2070,7 @@ async def test_true_substitutes_blocked_piece_for_text_only_scorer(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2087,7 +2083,7 @@ async def test_refusal_scorer_short_circuits_on_blocked_by_default(self): scorer = _MockRefusalScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2099,7 +2095,7 @@ async def test_refusal_scorer_evaluates_partial_content_when_flag_on(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2112,7 +2108,7 @@ async def test_no_substitute_when_no_partial_content(self): msg = Message(message_pieces=[_make_blocked_piece()]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2123,10 +2119,10 @@ async def test_normal_piece_unaffected_by_flag(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_normal_piece()]) - scores_off = await scorer.score_async(msg) + scores_off = await scorer.score_async(scorable=MessageScorable(message=msg)) scorer.scored_pieces.clear() scorer.score_blocked_content = True - scores_on = await scorer.score_async(msg) + scores_on = await scorer.score_async(scorable=MessageScorable(message=msg)) assert scores_off[0].score_value == scores_on[0].score_value @@ -2136,7 +2132,7 @@ async def test_mixed_pieces_only_blocked_substituted(self): msg = Message(message_pieces=[_make_normal_piece(), _make_blocked_piece(partial_content="partial harmful")]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg) + scores = await scorer.score_async(scorable=MessageScorable(message=msg)) assert len(scores) == 1 # TrueFalseScorer aggregates assert len(scorer.scored_pieces) == 2 @@ -2154,7 +2150,7 @@ async def test_skip_on_error_true_without_flag_skips_blocked(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert scores == [] async def test_skip_on_error_true_with_flag_does_not_skip_when_partial_content(self): @@ -2162,7 +2158,7 @@ async def test_skip_on_error_true_with_flag_does_not_skip_when_partial_content(s msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2171,7 +2167,7 @@ async def test_skip_on_error_true_with_flag_still_skips_when_no_partial_content( msg = Message(message_pieces=[_make_blocked_piece()]) scorer.score_blocked_content = True - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert scores == [] async def test_skip_on_error_skips_error_type_without_response_error_flag(self): @@ -2188,7 +2184,7 @@ async def test_skip_on_error_skips_error_type_without_response_error_flag(self): ] ) - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert scores == [] assert scorer.scored_pieces == [] @@ -2206,7 +2202,7 @@ async def test_skip_on_error_scores_structured_refusal_as_text(self, validator: piece = _make_blocked_piece(structured_refusal=refusal) msg = Message(message_pieces=[piece]) - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert len(scores) == 1 assert scorer.scored_pieces[0].id == piece.id @@ -2234,7 +2230,7 @@ async def test_skip_on_error_still_skips_mixed_structured_and_runtime_errors(sel ] ) - scores = await scorer.score_async(msg, skip_on_error_result=True) + scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) assert scores == [] assert scorer.scored_pieces == [] diff --git a/tests/unit/score/test_scorer_evaluator.py b/tests/unit/score/test_scorer_evaluator.py index a2d2016719..21eb1e6a53 100644 --- a/tests/unit/score/test_scorer_evaluator.py +++ b/tests/unit/score/test_scorer_evaluator.py @@ -6,6 +6,7 @@ import numpy as np import pytest +from pyrit.memory import MemoryInterface from pyrit.models import Message, MessagePiece from pyrit.score import ( FloatScaleScorer, @@ -26,8 +27,9 @@ @pytest.fixture def mock_harm_scorer(): scorer = MagicMock(spec=FloatScaleScorer) - scorer._memory = MagicMock() + scorer._memory = MagicMock(spec=MemoryInterface) scorer._memory.add_message_to_memory = MagicMock() + scorer._memory.get_message_pieces.return_value = [] # Create a mock identifier with a controllable hash property mock_identifier = MagicMock() mock_identifier.hash = "test_hash_456" @@ -40,8 +42,9 @@ def mock_harm_scorer(): @pytest.fixture def mock_objective_scorer(): scorer = MagicMock(spec=TrueFalseScorer) - scorer._memory = MagicMock() + scorer._memory = MagicMock(spec=MemoryInterface) scorer._memory.add_message_to_memory = MagicMock() + scorer._memory.get_message_pieces.return_value = [] # Create a mock identifier with a controllable hash property mock_identifier = MagicMock() mock_identifier.hash = "test_hash_123" diff --git a/tests/unit/score/test_self_ask_category.py b/tests/unit/score/test_self_ask_category.py index 5d6bc50d82..5da94f78f7 100644 --- a/tests/unit/score/test_self_ask_category.py +++ b/tests/unit/score/test_self_ask_category.py @@ -11,13 +11,8 @@ from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece -from pyrit.score import ( - ContentClassifier, - ContentClassifierCategory, - ContentClassifierPaths, - SelfAskCategoryScorer, -) +from pyrit.models import Message, MessagePiece, MessageScorable +from pyrit.score import ContentClassifier, ContentClassifierCategory, ContentClassifierPaths, SelfAskCategoryScorer HARM_CLASSIFIER = ContentClassifier.from_yaml(ContentClassifierPaths.HARMFUL_CONTENT_CLASSIFIER.value) @@ -298,7 +293,7 @@ async def test_blocked_response_returns_false_without_invoking_llm(patch_central ) blocked_message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(blocked_message) + scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) chat_target.send_prompt_async.assert_not_called() assert len(scores) == 1 diff --git a/tests/unit/score/test_self_ask_question_answer_scorer.py b/tests/unit/score/test_self_ask_question_answer_scorer.py index a8656ac480..a849338116 100644 --- a/tests/unit/score/test_self_ask_question_answer_scorer.py +++ b/tests/unit/score/test_self_ask_question_answer_scorer.py @@ -5,7 +5,7 @@ import pytest -from pyrit.models import ComponentIdentifier, MessagePiece, Score, UnvalidatedScore +from pyrit.models import ComponentIdentifier, MessagePiece, MessageScorable, Score, ScoringExpectation, UnvalidatedScore from pyrit.prompt_target import PromptTarget from pyrit.score.true_false.self_ask_question_answer_scorer import SelfAskQuestionAnswerScorer @@ -40,7 +40,9 @@ async def test_score_async_returns_score_from_unvalidated(mock_chat_target): "pyrit.score.true_false.self_ask_question_answer_scorer._run_llm_scoring_async", new=AsyncMock(return_value=unvalidated), ): - scores = await scorer.score_async(message, objective="2+2=?\nanswer: 4") + scores = await scorer.score_async( + scorable=MessageScorable(message=message), expectation=ScoringExpectation(objective="2+2=?\nanswer: 4") + ) assert len(scores) == 1 assert isinstance(scores[0], Score) diff --git a/tests/unit/score/test_self_ask_refusal.py b/tests/unit/score/test_self_ask_refusal.py index 4041eeaa64..217a025ac5 100644 --- a/tests/unit/score/test_self_ask_refusal.py +++ b/tests/unit/score/test_self_ask_refusal.py @@ -19,6 +19,7 @@ JsonResponseConfig, Message, MessagePiece, + MessageScorable, SeedPrompt, ) from pyrit.score import JsonSchemaResponseHandler, RefusalScorerPaths, SelfAskRefusalScorer @@ -254,7 +255,7 @@ async def test_score_async_filtered_response(patch_central_database): conversation_id=str(uuid4()), ).to_message() memory.add_message_pieces_to_memory(message_pieces=request.message_pieces) - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].score_value == "true" diff --git a/tests/unit/score/test_shieldgemma_scorer.py b/tests/unit/score/test_shieldgemma_scorer.py index cad3a54b2d..f365cc7ee3 100644 --- a/tests/unit/score/test_shieldgemma_scorer.py +++ b/tests/unit/score/test_shieldgemma_scorer.py @@ -9,7 +9,7 @@ from pyrit.exceptions import InvalidJsonException from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import JSON_SCHEMA_METADATA_KEY, Message, MessagePiece +from pyrit.models import JSON_SCHEMA_METADATA_KEY, Message, MessagePiece, MessageScorable from pyrit.prompt_target import PromptTarget from pyrit.score import ( ShieldGemmaGuideline, @@ -142,7 +142,7 @@ async def test_response_scoring_excludes_a_stored_user_turn(sqlite_instance: Mem target = _mock_target("No") scorer = ShieldGemmaScorer(chat_target=target, guideline=CUSTOM_GUIDELINE) - await scorer.score_async(response) + await scorer.score_async(scorable=MessageScorable(message=response)) sent = _sent_request(target) assert "Chatbot Response: A response judged on its own." in sent @@ -231,7 +231,7 @@ async def test_multiple_pieces_keep_every_verdict_and_report_the_aggregate( ) message.set_response_not_in_memory() - scores = await scorer.score_async(message) + scores = await scorer.score_async(scorable=MessageScorable(message=message)) assert target.send_prompt_async.call_count == 2 assert scores[0].get_value() is True diff --git a/tests/unit/score/test_substring.py b/tests/unit/score/test_substring.py index 64f6603481..ff014eead8 100644 --- a/tests/unit/score/test_substring.py +++ b/tests/unit/score/test_substring.py @@ -10,7 +10,7 @@ from pyrit.analytics import ApproximateTextMatching, ExactTextMatching from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece +from pyrit.models import MessagePiece, MessageScorable from pyrit.score import SubStringScorer @@ -27,7 +27,7 @@ async def test_score_async_unsupported_data_type_returns_false( scorer = SubStringScorer(substring="test", categories=["new_category"]) # With raise_on_no_valid_pieces=False (default), returns False for unsupported data types - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale diff --git a/tests/unit/score/test_true_false_composite_scorer.py b/tests/unit/score/test_true_false_composite_scorer.py index 5f9458161e..e3b73f1599 100644 --- a/tests/unit/score/test_true_false_composite_scorer.py +++ b/tests/unit/score/test_true_false_composite_scorer.py @@ -6,13 +6,8 @@ import pytest from pyrit.memory.central_memory import CentralMemory -from pyrit.models import ComponentIdentifier, MessagePiece, Score -from pyrit.score import ( - FloatScaleScorer, - TrueFalseCompositeScorer, - TrueFalseScoreAggregator, - TrueFalseScorer, -) +from pyrit.models import ComponentIdentifier, MessagePiece, MessageScorable, Score, ScoringExpectation +from pyrit.score import FloatScaleScorer, TrueFalseCompositeScorer, TrueFalseScoreAggregator, TrueFalseScorer def _mock_scorer_id(name: str = "MockScorer") -> ComponentIdentifier: @@ -82,7 +77,7 @@ def false_scorer(patch_central_database): async def test_composite_scorer_and_all_true(mock_request, true_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer, true_scorer]) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -92,7 +87,7 @@ async def test_composite_scorer_and_all_true(mock_request, true_scorer): async def test_composite_scorer_and_one_false(mock_request, true_scorer, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer, false_scorer]) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a false score" in scores[0].score_rationale @@ -102,7 +97,7 @@ async def test_composite_scorer_and_one_false(mock_request, true_scorer, false_s async def test_composite_scorer_or_all_false(mock_request, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.OR, scorers=[false_scorer, false_scorer]) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a false score" in scores[0].score_rationale @@ -112,7 +107,7 @@ async def test_composite_scorer_or_all_false(mock_request, false_scorer): async def test_composite_scorer_or_one_true(mock_request, true_scorer, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.OR, scorers=[true_scorer, false_scorer]) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -123,7 +118,7 @@ async def test_composite_scorer_majority_true(mock_request, true_scorer, false_s aggregator=TrueFalseScoreAggregator.MAJORITY, scorers=[true_scorer, true_scorer, false_scorer] ) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -138,7 +133,7 @@ async def test_composite_scorer_majority_false(mock_request, true_scorer, false_ aggregator=TrueFalseScoreAggregator.MAJORITY, scorers=[true_scorer, false_scorer, false_scorer] ) - scores = await scorer.score_async(mock_request) + scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a true score" in scores[0].score_rationale @@ -164,7 +159,9 @@ async def test_composite_scorer_with_task(mock_request, true_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer]) task = "test task" - scores = await scorer.score_async(mock_request, objective=task) + scores = await scorer.score_async( + scorable=MessageScorable(message=mock_request), expectation=ScoringExpectation(objective=task) + ) assert len(scores) == 1 assert scores[0].objective == task @@ -185,7 +182,7 @@ async def test_composite_scorer_raises_when_message_piece_id_is_none(true_scorer message = piece.to_message() with pytest.raises(RuntimeError, match="Message piece must have an ID"): - await scorer.score_async(message) + await scorer.score_async(scorable=MessageScorable(message=message)) def test_get_chat_target_returns_first_available(patch_central_database): diff --git a/tests/unit/score/test_true_false_inverter.py b/tests/unit/score/test_true_false_inverter.py index 28618d9e3a..7bd459f3da 100644 --- a/tests/unit/score/test_true_false_inverter.py +++ b/tests/unit/score/test_true_false_inverter.py @@ -9,7 +9,7 @@ from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece +from pyrit.models import MessagePiece, MessageScorable from pyrit.score import SubStringScorer, TrueFalseInverterScorer @@ -28,7 +28,7 @@ async def test_score_async_unsupported_data_type_inverts_false_to_true( # With raise_on_no_valid_pieces=False (default), the inner scorer returns False, # and the inverter inverts it to True - scores = await scorer.score_async(request) + scores = await scorer.score_async(scorable=MessageScorable(message=request)) assert len(scores) == 1 # Inverter inverts False -> True assert scores[0].get_value() is True From 47fbf14be4ba0c97e8002f2adf672d1701467a52 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 08:46:01 -0700 Subject: [PATCH 02/11] Remove scorer contract surface that Part A does not use Phase 1 shipped every scorable, condition, and scope type the design proposal names. Most of them have no consumer in this change and no phase in Part A (gist phases 1-6) that can construct or read one. The proposal itself says "nothing in Part A populates" a ScoringScope. Removed, with the phase that will reintroduce each: - SurfaceScorable and ScoringScope - phase 10, Part B, which is the first phase that populates a scope. - TraceScorable - phase 5, where it arrives with TraceSource and the Observation model it feeds. - ConversationScorable - never constructed in Part A. Phase 6 widens inside the scorer, as ConversationScorer reads conversation_id off the message today. - ToolCalled, ToolSequence, OutputMatches, Condition, and ScoringExpectation.conditions - phase 4 transports an expectation and phase 6 is the first reader. Until then a caller who sets conditions gets silence rather than an error, so ScoringExpectation holds objective alone. - ScoringExpectation.extra - named by no phase. - Volatility and the volatility ClassVars - nothing branches on them; their consumers are memoization in phase 3 and the lifecycle in phase 11. ScoringScope also carried a defect no consumer caught: window is declared tuple[datetime, datetime] | None and a test passed the string "last_turn". A frozen dataclass does not validate, so the test passed. Kept: MessageScorable, ContentScorable, MessageReferenceScorable, and ScoringExpectation. MessageReferenceScorable has a resolver and no producer, but phase 2 persists Score.scorable and a MessageScorable holding a live Message cannot round-trip through a column, so an id-based form is the persisted form. The unsupported-scorable test now uses a module-local dataclass instead of ConversationScorable, which proves the guard rejects anything outside _SUPPORTED_SCORABLES rather than one known sibling. No runtime behavior changes. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/models/__init__.py | 18 ----- pyrit/models/score/__init__.py | 23 +----- pyrit/models/score/expectation.py | 33 +-------- pyrit/models/score/scorable.py | 72 +----------------- pyrit/models/score/scoring_scope.py | 28 ------- pyrit/score/message_scorer.py | 2 +- pyrit/score/scorer.py | 4 +- tests/unit/models/test_scorable.py | 99 +++---------------------- tests/unit/score/test_message_scorer.py | 15 +++- 9 files changed, 28 insertions(+), 266 deletions(-) delete mode 100644 pyrit/models/score/scoring_scope.py diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 5293afab4c..10d248e593 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -79,23 +79,14 @@ from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT from pyrit.models.retry_event import RetryEvent from pyrit.models.score import ( - Condition, ContentScorable, - ConversationScorable, MessageReferenceScorable, MessageScorable, - OutputMatches, Scorable, Score, ScoreType, ScoringExpectation, - ScoringScope, - SurfaceScorable, - ToolCalled, - ToolSequence, - TraceScorable, UnvalidatedScore, - Volatility, ) # Seeds - import from new seeds submodule for forward compatibility @@ -151,14 +142,12 @@ "ComponentType", "compute_eval_hash", "config_hash", - "Condition", "ContentScorable", "ConverterIdentifier", "Conversation", "ConversationReference", "ConversationRetry", "ConversationRetryReason", - "ConversationScorable", "ConversationStats", "ConversationType", "construct_response_from_request", @@ -196,7 +185,6 @@ "Modality", "NextMessageSystemPromptPaths", "ObjectiveTargetEvaluationIdentifier", - "OutputMatches", "Parameter", "ParameterDestination", "PromptDataType", @@ -211,7 +199,6 @@ "Score", "ScoreType", "ScoringExpectation", - "ScoringScope", "ScenarioEvaluationIdentifier", "ScorerEvaluationIdentifier", "ScorerIdentifier", @@ -234,7 +221,6 @@ "sort_message_pieces", "StrategyResult", "StrategyResultT", - "SurfaceScorable", "TARGET_EVAL_PARAM_FALLBACKS", "TARGET_EVAL_PARAMS", "TargetCapabilities", @@ -243,11 +229,7 @@ "TOKEN_USAGE_METADATA_PREFIX", "TokenUsage", "ToolCall", - "ToolCalled", - "ToolSequence", - "TraceScorable", "UnvalidatedScore", - "Volatility", "read_usage_int", "read_usage_value", "validate_registry_name", diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py index 6bf0fdbe76..7f78b8ed4e 100644 --- a/pyrit/models/score/__init__.py +++ b/pyrit/models/score/__init__.py @@ -8,37 +8,18 @@ ``ScoringExpectation`` (what to look for) — and returns ``Score`` objects. """ -from pyrit.models.score.expectation import Condition, OutputMatches, ScoringExpectation, ToolCalled, ToolSequence -from pyrit.models.score.scorable import ( - ContentScorable, - ConversationScorable, - MessageReferenceScorable, - MessageScorable, - Scorable, - SurfaceScorable, - TraceScorable, - Volatility, -) +from pyrit.models.score.expectation import ScoringExpectation +from pyrit.models.score.scorable import ContentScorable, MessageReferenceScorable, MessageScorable, Scorable from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore -from pyrit.models.score.scoring_scope import ScoringScope __all__ = [ "ComponentIdentifierField", - "Condition", "ContentScorable", - "ConversationScorable", "MessageReferenceScorable", "MessageScorable", - "OutputMatches", "Scorable", "Score", "ScoreType", "ScoringExpectation", - "ScoringScope", - "SurfaceScorable", - "ToolCalled", - "ToolSequence", - "TraceScorable", "UnvalidatedScore", - "Volatility", ] diff --git a/pyrit/models/score/expectation.py b/pyrit/models/score/expectation.py index 0a89c3c080..924f52d487 100644 --- a/pyrit/models/score/expectation.py +++ b/pyrit/models/score/expectation.py @@ -3,32 +3,7 @@ from __future__ import annotations -from dataclasses import dataclass, field - - -@dataclass(frozen=True) -class ToolCalled: - """The named tool appears in the evidence.""" - - name: str - arguments: dict[str, str] | None = None - - -@dataclass(frozen=True) -class ToolSequence: - """The named tools appear, in this relative order.""" - - names: tuple[str, ...] - - -@dataclass(frozen=True) -class OutputMatches: - """The response contains or equals a ground-truth answer.""" - - value: str - - -Condition = ToolCalled | ToolSequence | OutputMatches +from dataclasses import dataclass @dataclass(frozen=True) @@ -39,12 +14,6 @@ class ScoringExpectation: An expectation is a single parameter that attacks forward without inspecting it, so a question authored in a technique configuration or a seed can reach a scorer through an attack that knows nothing about it. - - Conditions are neutral about polarity: a condition says what to detect, never - whether detecting it is good or bad. Wrap a scorer in ``TrueFalseInverterScorer`` - to express the negative case. """ objective: str | None = None - conditions: tuple[Condition, ...] = () - extra: dict[str, str] = field(default_factory=dict) diff --git a/pyrit/models/score/scorable.py b/pyrit/models/score/scorable.py index 5864a22840..cba4a4222d 100644 --- a/pyrit/models/score/scorable.py +++ b/pyrit/models/score/scorable.py @@ -4,41 +4,13 @@ from __future__ import annotations from dataclasses import dataclass -from enum import Enum -from typing import TYPE_CHECKING, ClassVar +from typing import TYPE_CHECKING if TYPE_CHECKING: import uuid from pyrit.models.literals import ChatMessageRole, PromptDataType from pyrit.models.messages import Message - from pyrit.models.score.scoring_scope import ScoringScope - - -class Volatility(Enum): - """Whether two reads of the same scorable can disagree.""" - - STABLE = "stable" - APPEND_ONLY = "append_only" - MUTABLE = "mutable" - - -@dataclass(frozen=True) -class ConversationScorable: - """ - A stored conversation, resolved from memory. - - A conversation is append-only rather than stable because it grows while the attack - runs; reading it after turn one and after turn five gives different evidence. Leave - ``through_sequence`` unset to mean "as far as it goes right now"; set it to name a - reproducible prefix. - """ - - volatility: ClassVar[Volatility] = Volatility.APPEND_ONLY - - conversation_id: str - through_sequence: int | None = None - role_filter: ChatMessageRole | None = None @dataclass(frozen=True) @@ -51,8 +23,6 @@ class MessageScorable: ``MessageReferenceScorable`` when only the piece ids are known. """ - volatility: ClassVar[Volatility] = Volatility.STABLE - message: Message role_filter: ChatMessageRole | None = None skip_on_error_result: bool = False @@ -67,8 +37,6 @@ class MessageReferenceScorable: ``MessageScorable`` when the caller already holds the message. """ - volatility: ClassVar[Volatility] = Volatility.STABLE - message_piece_ids: tuple[uuid.UUID | str, ...] role_filter: ChatMessageRole | None = None skip_on_error_result: bool = False @@ -78,44 +46,8 @@ class MessageReferenceScorable: class ContentScorable: """Loose content with no conversation behind it.""" - volatility: ClassVar[Volatility] = Volatility.STABLE - value: str data_type: PromptDataType = "text" -@dataclass(frozen=True) -class SurfaceScorable: - """A location that may or may not have been written.""" - - volatility: ClassVar[Volatility] = Volatility.MUTABLE - - uri: str - surface: str = "file" - scope: ScoringScope | None = None - - -@dataclass(frozen=True) -class TraceScorable: - """ - A recorded agent run, or named steps within it. - - A span is the tracing term for one recorded step — a single tool call, model - request, or nested operation. A trace is the tree of spans for one run. - """ - - volatility: ClassVar[Volatility] = Volatility.APPEND_ONLY - - trace_ids: tuple[str, ...] = () - span_name: str | None = None - scope: ScoringScope | None = None - - -Scorable = ( - ConversationScorable - | MessageScorable - | MessageReferenceScorable - | ContentScorable - | SurfaceScorable - | TraceScorable -) +Scorable = MessageScorable | MessageReferenceScorable | ContentScorable diff --git a/pyrit/models/score/scoring_scope.py b/pyrit/models/score/scoring_scope.py deleted file mode 100644 index 8247c5612f..0000000000 --- a/pyrit/models/score/scoring_scope.py +++ /dev/null @@ -1,28 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -from __future__ import annotations - -from dataclasses import dataclass, field -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from datetime import datetime - - -@dataclass(frozen=True) -class ScoringScope: - """ - Bounds a scorable that names a location rather than a stored record. - - A surface URI or a trace id names a place, not a thing, so it cannot say on its - own which run produced what is there now. A scope narrows that question with a - time window and correlation labels. - - Labels are the extension point for frameworks built on PyRIT that control how - evidence is emitted and can therefore supply stronger correlation keys. PyRIT - never interprets them. - """ - - window: tuple[datetime, datetime] | None = None - labels: dict[str, str] = field(default_factory=dict) diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index 80787a45fb..d65eda266e 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -212,7 +212,7 @@ def _resolve_message(self, scorable: MessageScorable | MessageReferenceScorable if len(messages) != 1: raise ValueError( f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " - "Reference pieces from a single message, or use a ConversationScorable." + "Reference pieces from a single message." ) return messages[0] diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 1a1c07df9c..9353989b8f 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -318,8 +318,8 @@ def _resolve_score_inputs( ) if scorable is None: - # A supplied message maps to MessageScorable, never ConversationScorable: the - # caller asked about this message, not about everything stored alongside it. + # The caller asked about this message, not about everything stored alongside it, + # so the shim maps to the exact message rather than widening to its conversation. scorable = MessageScorable( message=cast("Message", message), role_filter=role_filter, diff --git a/tests/unit/models/test_scorable.py b/tests/unit/models/test_scorable.py index c8322ca4c7..aa858a1a8d 100644 --- a/tests/unit/models/test_scorable.py +++ b/tests/unit/models/test_scorable.py @@ -6,21 +6,7 @@ import pytest -from pyrit.models import ( - ContentScorable, - ConversationScorable, - MessagePiece, - MessageReferenceScorable, - MessageScorable, - OutputMatches, - ScoringExpectation, - ScoringScope, - SurfaceScorable, - ToolCalled, - ToolSequence, - TraceScorable, - Volatility, -) +from pyrit.models import ContentScorable, MessagePiece, MessageReferenceScorable, MessageScorable, ScoringExpectation def _message(): @@ -28,40 +14,16 @@ def _message(): @pytest.mark.parametrize( - "scorable, expected_volatility", - [ - (ConversationScorable(conversation_id="c1"), Volatility.APPEND_ONLY), - (MessageScorable(message=_message()), Volatility.STABLE), - (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), Volatility.STABLE), - (ContentScorable(value="hello"), Volatility.STABLE), - (SurfaceScorable(uri="/tmp/out.txt"), Volatility.MUTABLE), - (TraceScorable(trace_ids=("t1",)), Volatility.APPEND_ONLY), - ], -) -def test_scorable_volatility(scorable, expected_volatility): - assert scorable.volatility is expected_volatility - - -@pytest.mark.parametrize( - "scorable", + "scorable, field_name", [ - ConversationScorable(conversation_id="c1"), - MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), - ContentScorable(value="hello"), - SurfaceScorable(uri="/tmp/out.txt"), - TraceScorable(trace_ids=("t1",)), + (MessageScorable(message=_message()), "message"), + (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), + (ContentScorable(value="hello"), "value"), ], ) -def test_scorable_is_frozen(scorable): +def test_scorable_is_frozen(scorable, field_name): with pytest.raises(dataclasses.FrozenInstanceError): - scorable.volatility = Volatility.MUTABLE - - -def test_conversation_scorable_defaults(): - scorable = ConversationScorable(conversation_id="c1") - - assert scorable.through_sequence is None - assert scorable.role_filter is None + setattr(scorable, field_name, "changed") def test_message_scorable_carries_the_message(): @@ -86,42 +48,8 @@ def test_content_scorable_defaults_to_text(): assert ContentScorable(value="hello").data_type == "text" -def test_surface_scorable_defaults_to_file(): - scorable = SurfaceScorable(uri="/tmp/out.txt") - - assert scorable.surface == "file" - assert scorable.scope is None - - -def test_trace_scorable_accepts_a_scope(): - scope = ScoringScope(window="last_turn") - scorable = TraceScorable(trace_ids=("t1",), span_name="tool_call", scope=scope) - - assert scorable.scope is scope - assert scorable.span_name == "tool_call" - - -def test_scoring_scope_defaults(): - scope = ScoringScope() - - assert scope.window is None - assert scope.labels == {} - - def test_expectation_defaults(): - expectation = ScoringExpectation() - - assert expectation.objective is None - assert expectation.conditions == () - assert expectation.extra == {} - - -def test_expectation_carries_conditions(): - conditions = (ToolCalled(name="send_email"), ToolSequence(names=("a", "b")), OutputMatches(value="42")) - expectation = ScoringExpectation(objective="exfiltrate", conditions=conditions) - - assert expectation.objective == "exfiltrate" - assert expectation.conditions == conditions + assert ScoringExpectation().objective is None def test_expectation_is_frozen(): @@ -131,14 +59,5 @@ def test_expectation_is_frozen(): expectation.objective = "something else" -def test_tool_called_defaults(): - condition = ToolCalled(name="send_email") - - assert condition.name == "send_email" - assert condition.arguments is None - - def test_expectations_with_equal_values_compare_equal(): - assert ScoringExpectation(objective="a", conditions=(ToolCalled(name="t"),)) == ScoringExpectation( - objective="a", conditions=(ToolCalled(name="t"),) - ) + assert ScoringExpectation(objective="a") == ScoringExpectation(objective="a") diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py index 0655a84c33..93d532201f 100644 --- a/tests/unit/score/test_message_scorer.py +++ b/tests/unit/score/test_message_scorer.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import dataclasses import uuid import pytest @@ -9,7 +10,6 @@ from pyrit.models import ( ComponentIdentifier, ContentScorable, - ConversationScorable, Message, MessagePiece, MessageReferenceScorable, @@ -21,6 +21,13 @@ from pyrit.score.message_scorer import extract_objective_from_previous_turn +@dataclasses.dataclass(frozen=True) +class UnsupportedScorable: + """A scorable kind no message scorer handles.""" + + uri: str + + class PermissiveValidator(ScorerPromptValidator): def validate(self, message, objective=None): pass @@ -135,8 +142,8 @@ async def test_content_scorable_is_never_persisted(self): async def test_unsupported_scorable_raises_type_error(self): scorer = RecordingScorer() - with pytest.raises(TypeError, match="cannot score ConversationScorable"): - await scorer.score_async(scorable=ConversationScorable(conversation_id=str(uuid.uuid4()))) + with pytest.raises(TypeError, match="cannot score UnsupportedScorable"): + await scorer.score_async(scorable=UnsupportedScorable(uri="/tmp/out.txt")) # type: ignore[arg-type] @pytest.mark.usefixtures("patch_central_database") @@ -234,7 +241,7 @@ async def test_keyword_message_maps_to_message_scorable(self): assert scorer.scored_messages == [message] async def test_message_does_not_widen_to_the_stored_conversation(self, sqlite_instance: MemoryInterface): - """A supplied message must map to MessageScorable, never ConversationScorable.""" + """The shim scores the supplied message, never the whole conversation behind it.""" conversation_id = str(uuid.uuid4()) sqlite_instance.add_message_to_memory( request=MessagePiece( From 4788a643b68e27e1e04c4f31e4f6a5b9c69ca0e7 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 10:08:09 -0700 Subject: [PATCH 03/11] Move scorables into pyrit.score and let them resolve themselves MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit MessageScorer did the resolving: an isinstance chain over a union of dumb records, plus a memory read and getattr duck-typing for the filter fields. That put data plumbing in the scorer and did not extend to the trace and surface scorables the design adds later. Scorables now resolve themselves. Scorable becomes an abstract root type instead of a union, SingleMessageScorable declares role_filter, skip_on_error_result, and an abstract resolve_message, and each concrete scorable says how it resolves: MessageScorable returns what it holds, MessageReferenceScorable reads memory, ContentScorable builds the never-persisted message. Adding a scorable is now adding a class, with no union to edit. Resolution needs pyrit.memory, which models.instructions.md does not allow on a model, so the scorables move from pyrit/models/score/ to pyrit/score/. That also clears a real cost: MessageScorable holds a Message but could not import one at runtime, because pyrit.models.messages imports ComponentIdentifierField back out of pyrit.models.score. ScoringExpectation stays in pyrit.models — it has no behavior, and a later phase has SeedGroup produce one. Only identity resolution belongs on a scorable. A scorer that derives a different kind of evidence from the same reference, such as trace ids or a file write, still does that widening itself. Removes _SUPPORTED_SCORABLES, the three-arm _resolve_message, and both getattr(scorable, ...) calls. This also carries three fixes from the review of the previous commit: - Scorer no longer defines the message hooks. _score_async, _score_piece_async, and _get_supported_pieces move to MessageScorer, and Scorer._score_scorable_async becomes abstract. A scorer deriving directly from Scorer that implemented only _score_piece_async used to build fine and then fail at score time with a confusing TypeError; it now fails at instantiation with a clear abstract-method error. - A partially resolvable MessageReferenceScorable raises and names only the ids that are missing, rather than only raising when nothing resolved. - The scorable module no longer needs a TYPE_CHECKING workaround for Message, since the circular import it worked around does not reach pyrit.score. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/executor/attack/multi_turn/crescendo.py | 10 +- .../executor/attack/multi_turn/red_teaming.py | 2 +- pyrit/models/__init__.py | 15 +- pyrit/models/score/__init__.py | 10 +- pyrit/models/score/scorable.py | 53 ------ pyrit/score/__init__.py | 12 ++ pyrit/score/audio_transcript_scorer.py | 3 +- pyrit/score/conversation_scorer.py | 9 +- pyrit/score/message_scorer.py | 108 ++++++------ pyrit/score/scorable.py | 147 ++++++++++++++++ pyrit/score/scorer.py | 52 +----- .../float_scale_threshold_scorer.py | 11 +- .../true_false/true_false_composite_scorer.py | 11 +- .../true_false/true_false_inverter_scorer.py | 11 +- .../test_scorer_contract.py | 22 ++- .../attack/multi_turn/test_tree_of_attacks.py | 3 +- tests/unit/models/test_expectation.py | 23 +++ tests/unit/models/test_scorable.py | 63 ------- tests/unit/score/test_azure_content_filter.py | 3 +- .../score/test_conversation_history_scorer.py | 14 +- .../test_float_scale_threshold_scorer.py | 4 +- tests/unit/score/test_gandalf_scorer.py | 4 +- tests/unit/score/test_insecure_code_scorer.py | 12 +- tests/unit/score/test_message_scorer.py | 47 ++++- tests/unit/score/test_plagiarism_scorer.py | 4 +- .../unit/score/test_question_answer_scorer.py | 4 +- tests/unit/score/test_scorable.py | 160 ++++++++++++++++++ tests/unit/score/test_scorer.py | 3 +- tests/unit/score/test_self_ask_category.py | 10 +- .../test_self_ask_question_answer_scorer.py | 3 +- tests/unit/score/test_self_ask_refusal.py | 3 +- tests/unit/score/test_shieldgemma_scorer.py | 3 +- tests/unit/score/test_substring.py | 4 +- .../score/test_true_false_composite_scorer.py | 10 +- tests/unit/score/test_true_false_inverter.py | 4 +- 35 files changed, 532 insertions(+), 325 deletions(-) delete mode 100644 pyrit/models/score/scorable.py create mode 100644 pyrit/score/scorable.py create mode 100644 tests/unit/models/test_expectation.py delete mode 100644 tests/unit/models/test_scorable.py create mode 100644 tests/unit/score/test_scorable.py diff --git a/pyrit/executor/attack/multi_turn/crescendo.py b/pyrit/executor/attack/multi_turn/crescendo.py index 7b28bd773d..9fb48abfec 100644 --- a/pyrit/executor/attack/multi_turn/crescendo.py +++ b/pyrit/executor/attack/multi_turn/crescendo.py @@ -31,14 +31,20 @@ ConversationType, Message, MessagePiece, - MessageScorable, Score, ScoringExpectation, SeedPrompt, ) from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName, TargetRequirements -from pyrit.score import FloatScaleThresholdScorer, NumericRubric, Scorer, SelfAskRefusalScorer, SelfAskScaleScorer +from pyrit.score import ( + FloatScaleThresholdScorer, + MessageScorable, + NumericRubric, + Scorer, + SelfAskRefusalScorer, + SelfAskScaleScorer, +) from pyrit.score.score_utils import normalize_score_to_float if TYPE_CHECKING: diff --git a/pyrit/executor/attack/multi_turn/red_teaming.py b/pyrit/executor/attack/multi_turn/red_teaming.py index 59fc6627e8..6ae5a19055 100644 --- a/pyrit/executor/attack/multi_turn/red_teaming.py +++ b/pyrit/executor/attack/multi_turn/red_teaming.py @@ -33,13 +33,13 @@ ConversationReference, ConversationType, Message, - MessageScorable, Score, ScoringExpectation, ) from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName from pyrit.prompt_target.common.target_requirements import TargetRequirements +from pyrit.score import MessageScorable if TYPE_CHECKING: from collections.abc import Callable diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 10d248e593..2cbbcac1c2 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -78,16 +78,7 @@ from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT from pyrit.models.retry_event import RetryEvent -from pyrit.models.score import ( - ContentScorable, - MessageReferenceScorable, - MessageScorable, - Scorable, - Score, - ScoreType, - ScoringExpectation, - UnvalidatedScore, -) +from pyrit.models.score import Score, ScoreType, ScoringExpectation, UnvalidatedScore # Seeds - import from new seeds submodule for forward compatibility # Also keep imports from old locations for backward compatibility @@ -142,7 +133,6 @@ "ComponentType", "compute_eval_hash", "config_hash", - "ContentScorable", "ConverterIdentifier", "Conversation", "ConversationReference", @@ -180,8 +170,6 @@ "MEDIA_PATH_DATA_TYPES", "Message", "MessagePiece", - "MessageReferenceScorable", - "MessageScorable", "Modality", "NextMessageSystemPromptPaths", "ObjectiveTargetEvaluationIdentifier", @@ -195,7 +183,6 @@ "QuestionChoice", "REGISTRY_NAME_PATTERN", "ScaleDescription", - "Scorable", "Score", "ScoreType", "ScoringExpectation", diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py index 7f78b8ed4e..bc357c13c1 100644 --- a/pyrit/models/score/__init__.py +++ b/pyrit/models/score/__init__.py @@ -2,22 +2,18 @@ # Licensed under the MIT license. """ -Score types: what is scored, what it is scored against, and the result. +Score types: what a scorer is scored against, and the result. A scorer takes two inputs — a ``Scorable`` (what to look at) and a -``ScoringExpectation`` (what to look for) — and returns ``Score`` objects. +``ScoringExpectation`` (what to look for) — and returns ``Score`` objects. Scorables +resolve themselves against memory, so they live in ``pyrit.score`` rather than here. """ from pyrit.models.score.expectation import ScoringExpectation -from pyrit.models.score.scorable import ContentScorable, MessageReferenceScorable, MessageScorable, Scorable from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore __all__ = [ "ComponentIdentifierField", - "ContentScorable", - "MessageReferenceScorable", - "MessageScorable", - "Scorable", "Score", "ScoreType", "ScoringExpectation", diff --git a/pyrit/models/score/scorable.py b/pyrit/models/score/scorable.py deleted file mode 100644 index cba4a4222d..0000000000 --- a/pyrit/models/score/scorable.py +++ /dev/null @@ -1,53 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -from __future__ import annotations - -from dataclasses import dataclass -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - import uuid - - from pyrit.models.literals import ChatMessageRole, PromptDataType - from pyrit.models.messages import Message - - -@dataclass(frozen=True) -class MessageScorable: - """ - A message the caller already holds. - - Use this when the message is in hand — an attack scoring the response it just - received, or a caller scoring a message whose pieces were never persisted. Use - ``MessageReferenceScorable`` when only the piece ids are known. - """ - - message: Message - role_filter: ChatMessageRole | None = None - skip_on_error_result: bool = False - - -@dataclass(frozen=True) -class MessageReferenceScorable: - """ - Specific message pieces, resolved from memory by id. - - Use this when the pieces are persisted and only their ids are known. Use - ``MessageScorable`` when the caller already holds the message. - """ - - message_piece_ids: tuple[uuid.UUID | str, ...] - role_filter: ChatMessageRole | None = None - skip_on_error_result: bool = False - - -@dataclass(frozen=True) -class ContentScorable: - """Loose content with no conversation behind it.""" - - value: str - data_type: PromptDataType = "text" - - -Scorable = MessageScorable | MessageReferenceScorable | ContentScorable diff --git a/pyrit/score/__init__.py b/pyrit/score/__init__.py index 24f449e27c..1cabd5b035 100644 --- a/pyrit/score/__init__.py +++ b/pyrit/score/__init__.py @@ -33,6 +33,13 @@ from pyrit.score.float_scale.self_ask_scale_scorer import SelfAskScaleScorer, render_scale_system_prompt from pyrit.score.message_scorer import MessageScorer from pyrit.score.response_handler import CallableResponseHandler, JsonSchemaResponseHandler, ResponseHandler +from pyrit.score.scorable import ( + ContentScorable, + MessageReferenceScorable, + MessageScorable, + Scorable, + SingleMessageScorable, +) from pyrit.score.scorer import Scorer from pyrit.score.scorer_evaluation.metrics_type import MetricsType, RegistryUpdateBehavior from pyrit.score.scorer_evaluation.scorer_metrics import ( @@ -159,6 +166,7 @@ def __getattr__(name: str) -> object: "ContentClassifier", "ContentClassifierCategory", "ContentClassifierPaths", + "ContentScorable", "ConversationScorer", "CredentialLeakScorer", "DecodingScorer", @@ -188,6 +196,8 @@ def __getattr__(name: str) -> object: "LlamaGuardPolicy", "LlamaGuardScorer", "MarkdownInjectionScorer", + "MessageReferenceScorable", + "MessageScorable", "MessageScorer", "MethKeywordScorer", "MetricsType", @@ -215,6 +225,7 @@ def __getattr__(name: str) -> object: "render_shieldgemma_prompt", "render_true_false_system_prompt", "ResponseHandler", + "Scorable", "Scorer", "ScorerEvalDatasetFiles", "ScorerEvaluator", @@ -234,6 +245,7 @@ def __getattr__(name: str) -> object: "SelfAskRefusalScorer", "SelfAskScaleScorer", "SelfAskTrueFalseScorer", + "SingleMessageScorable", "ScorerPrinter", "SHIELDGEMMA_DEFAULT_POLICY_PATH", "ShieldGemmaGuideline", diff --git a/pyrit/score/audio_transcript_scorer.py b/pyrit/score/audio_transcript_scorer.py index 1977f52372..7ed29d505e 100644 --- a/pyrit/score/audio_transcript_scorer.py +++ b/pyrit/score/audio_transcript_scorer.py @@ -11,7 +11,8 @@ from pyrit.converter import AzureSpeechAudioToTextConverter from pyrit.memory import CentralMemory -from pyrit.models import MessagePiece, MessageScorable, Score, ScoringExpectation +from pyrit.models import MessagePiece, Score, ScoringExpectation +from pyrit.score.scorable import MessageScorable from pyrit.score.scorer import Scorer logger = logging.getLogger(__name__) diff --git a/pyrit/score/conversation_scorer.py b/pyrit/score/conversation_scorer.py index 1a78e88793..c3cb68230f 100644 --- a/pyrit/score/conversation_scorer.py +++ b/pyrit/score/conversation_scorer.py @@ -139,7 +139,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st raise NotImplementedError("ConversationScorer uses _score_async, not _score_piece_async") @abstractmethod - def _get_wrapped_scorer(self) -> Scorer: + def _get_wrapped_scorer(self) -> MessageScorer: """ Abstract method to enforce that ConversationScorer cannot be instantiated directly. @@ -200,6 +200,9 @@ def create_conversation_scorer( f"Scorer must be an instance of FloatScaleScorer or TrueFalseScorer." ) + # Both branches above narrow to a MessageScorer, which is what supplies _score_async. + wrapped_scorer: MessageScorer = scorer + # Dynamically create a class that inherits from both ConversationScorer and the scorer's base class class DynamicConversationScorer(ConversationScorer, scorer_base_class): # type: ignore[valid-type] # type: ignore[ty:unsupported-base] """Dynamic ConversationScorer that inherits from both ConversationScorer and the wrapped scorer's base class.""" @@ -207,9 +210,9 @@ class DynamicConversationScorer(ConversationScorer, scorer_base_class): # type: def __init__(self) -> None: # Initialize with the validator and wrapped scorer Scorer.__init__(self, validator=validator or ConversationScorer._DEFAULT_VALIDATOR) - self._wrapped_scorer = scorer + self._wrapped_scorer = wrapped_scorer - def _get_wrapped_scorer(self) -> Scorer: + def _get_wrapped_scorer(self) -> MessageScorer: """Return the wrapped scorer.""" return self._wrapped_scorer diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index d65eda266e..fe8124b19e 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -3,31 +3,21 @@ from __future__ import annotations +import asyncio import logging +from abc import abstractmethod from typing import TYPE_CHECKING from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException -from pyrit.models import ( - ContentScorable, - Message, - MessagePiece, - MessageReferenceScorable, - MessageScorable, - Scorable, - Score, - ScoringExpectation, - group_message_pieces_into_conversations, -) +from pyrit.score.scorable import Scorable, SingleMessageScorable from pyrit.score.scorer import Scorer if TYPE_CHECKING: from pyrit.memory import MemoryInterface + from pyrit.models import Message, MessagePiece, Score, ScoringExpectation logger = logging.getLogger(__name__) -#: Scorable kinds a MessageScorer can reduce to a single Message. -_SUPPORTED_SCORABLES = (MessageScorable, MessageReferenceScorable, ContentScorable) - def extract_objective_from_previous_turn(*, message: Message, memory: MemoryInterface) -> str: """ @@ -70,10 +60,10 @@ class MessageScorer(Scorer): """ Base class for scorers whose evidence is a single message. - Every message-shaped concern lives here: resolving a message scorable to a ``Message``, - substituting refusal and blocked content, validating pieces, applying the role and error - filters, and falling back to a neutral score. ``Scorer`` stays agnostic about what a - scorable is, so scorers over other kinds of evidence can sit beside this one. + Every message-shaped concern lives here: substituting refusal and blocked content, + validating pieces, applying the role and error filters, and falling back to a neutral + score. ``Scorer`` stays agnostic about what a scorable is, so scorers over other kinds of + evidence can sit beside this one. The scorable resolves itself to a ``Message``. Subclasses implement ``_score_async``, which still receives a ``Message``. """ @@ -89,8 +79,7 @@ async def _score_scorable_async( Resolve a message scorable and score the message it names. Args: - scorable (Scorable): A ``MessageScorable``, ``MessageReferenceScorable``, or - ``ContentScorable``. + scorable (Scorable): Any ``SingleMessageScorable``. expectation (ScoringExpectation | None): What to look for. infer_objective_from_request (bool): Deprecated; read the objective from the previous turn when the expectation carries none. @@ -105,15 +94,15 @@ async def _score_scorable_async( PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). """ - if not isinstance(scorable, _SUPPORTED_SCORABLES): + if not isinstance(scorable, SingleMessageScorable): raise TypeError( f"{self.__class__.__name__} scores messages, so it cannot score {type(scorable).__name__}. " "Pass a MessageScorable, a MessageReferenceScorable, or a ContentScorable." ) - message = self._resolve_message(scorable) - role_filter = getattr(scorable, "role_filter", None) - skip_on_error_result = getattr(scorable, "skip_on_error_result", False) + message = scorable.resolve_message(memory=self._memory) + role_filter = scorable.role_filter + skip_on_error_result = scorable.skip_on_error_result objective = expectation.objective if expectation else None # Structured refusals are persisted as blocked error pieces, but scorers should @@ -178,43 +167,52 @@ async def _score_scorable_async( return scores - def _resolve_message(self, scorable: MessageScorable | MessageReferenceScorable | ContentScorable) -> Message: + async def _score_async(self, message: Message, *, objective: str | None = None) -> list[Score]: """ - Return the message a message-shaped scorable names. + Score the given request response asynchronously. + + This default implementation scores all supported pieces in the message + and returns a flattened list of scores. Subclasses can override this method + to implement custom scoring logic (e.g., aggregating scores). - Loose content has no conversation behind it, so it becomes a message that is marked - as never persisted. Phase 2 gives scores a scorable of their own and removes this. + Args: + message (Message): The message to score. + objective (str | None): The objective to evaluate against. Defaults to None. Returns: - Message: The message to score. + list[Score]: A list of Score objects. + """ + if not message.message_pieces: + return [] - Raises: - ValueError: If the referenced pieces are not in memory or do not form one message. + # Score only the supported pieces + supported_pieces = self._get_supported_pieces(message) + + tasks = [self._score_piece_async(message_piece=piece, objective=objective) for piece in supported_pieces] + + if not tasks: + return [] + + # Run all piece-level scorings concurrently + piece_score_lists = await asyncio.gather(*tasks) + + # Flatten list[list[Score]] -> list[Score] + return [score for sublist in piece_score_lists for score in sublist] + + @abstractmethod + async def _score_piece_async(self, message_piece: MessagePiece, *, objective: str | None = None) -> list[Score]: + raise NotImplementedError + + def _get_supported_pieces(self, message: Message) -> list[MessagePiece]: """ - if isinstance(scorable, MessageScorable): - return scorable.message - - if isinstance(scorable, ContentScorable): - piece = MessagePiece( - role="user", - original_value=scorable.value, - original_value_data_type=scorable.data_type, - ) - piece.not_in_memory = True - return Message(message_pieces=[piece]) - - pieces = self._memory.get_message_pieces(prompt_ids=list(scorable.message_piece_ids)) - if not pieces: - raise ValueError(f"No message pieces found in memory for ids {list(scorable.message_piece_ids)}.") - - conversations = group_message_pieces_into_conversations(pieces) - messages = [message for conversation in conversations for message in conversation] - if len(messages) != 1: - raise ValueError( - f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " - "Reference pieces from a single message." - ) - return messages[0] + Get a list of supported message pieces for this scorer. + + Returns: + list[MessagePiece]: List of message pieces that are supported by this scorer's validator. + """ + return [ + piece for piece in message.message_pieces if self._validator.is_message_piece_supported(message_piece=piece) + ] def _should_skip_on_error(self, message: Message) -> bool: """ diff --git a/pyrit/score/scorable.py b/pyrit/score/scorable.py new file mode 100644 index 0000000000..48da9408c9 --- /dev/null +++ b/pyrit/score/scorable.py @@ -0,0 +1,147 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import uuid # noqa: TC003 (runtime-required by dataclass field annotations) +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from pyrit.models import ( # noqa: TC001 (runtime-required by dataclass field annotations) + ChatMessageRole, + Message, + MessagePiece, + PromptDataType, + group_message_pieces_into_conversations, +) + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +class Scorable(ABC): # noqa: B024 root type; each scorer family declares its own resolution + """ + What a scorer looks at. + + A scorable is a reference rather than a payload: it names the evidence instead of + carrying it. Each scorer accepts the scorable kinds it supports and rejects the rest. + """ + + +@dataclass(frozen=True, kw_only=True) +class SingleMessageScorable(Scorable, ABC): + """ + A scorable that names exactly one message. + + Subclasses resolve themselves. Only identity resolution belongs here — returning the + thing the scorable names. A scorer that derives a different kind of evidence from the + same reference, such as trace ids or a file write, does that widening itself. + """ + + role_filter: ChatMessageRole | None = None + skip_on_error_result: bool = False + + @abstractmethod + def resolve_message(self, *, memory: MemoryInterface) -> Message: + """ + Return the message this scorable names. + + Args: + memory (MemoryInterface): Memory to resolve references against. + + Returns: + Message: The message to score. + """ + + +@dataclass(frozen=True, kw_only=True) +class MessageScorable(SingleMessageScorable): + """ + A message the caller already holds. + + Use this when the message is in hand — an attack scoring the response it just + received, or a caller scoring a message whose pieces were never persisted. Use + ``MessageReferenceScorable`` when only the piece ids are known. + """ + + message: Message + + def resolve_message(self, *, memory: MemoryInterface) -> Message: + """ + Return the message the caller supplied. + + Args: + memory (MemoryInterface): Unused; the message is already in hand. + + Returns: + Message: The message to score. + """ + return self.message + + +@dataclass(frozen=True, kw_only=True) +class MessageReferenceScorable(SingleMessageScorable): + """ + Specific message pieces, resolved from memory by id. + + Use this when the pieces are persisted and only their ids are known. Use + ``MessageScorable`` when the caller already holds the message. + """ + + message_piece_ids: tuple[uuid.UUID | str, ...] + + def resolve_message(self, *, memory: MemoryInterface) -> Message: + """ + Return the single message the referenced pieces form. + + Args: + memory (MemoryInterface): Memory holding the pieces. + + Returns: + Message: The message to score. + + Raises: + ValueError: If any referenced piece is not in memory, or the pieces do not form + exactly one message. + """ + pieces = memory.get_message_pieces(prompt_ids=list(self.message_piece_ids)) + found = {str(piece.id) for piece in pieces} + missing = [str(piece_id) for piece_id in self.message_piece_ids if str(piece_id) not in found] + if missing: + raise ValueError(f"No message pieces found in memory for ids {missing}.") + + conversations = group_message_pieces_into_conversations(pieces) + messages = [message for conversation in conversations for message in conversation] + if len(messages) != 1: + raise ValueError( + f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " + "Reference pieces from a single message." + ) + return messages[0] + + +@dataclass(frozen=True, kw_only=True) +class ContentScorable(SingleMessageScorable): + """Loose content with no conversation behind it.""" + + value: str + data_type: PromptDataType = "text" + + def resolve_message(self, *, memory: MemoryInterface) -> Message: + """ + Return a message holding the loose content, marked as never persisted. + + Args: + memory (MemoryInterface): Unused; loose content has no conversation to read. + + Returns: + Message: The message to score. + """ + piece = MessagePiece( + role="user", + original_value=self.value, + original_value_data_type=self.data_type, + ) + piece.not_in_memory = True + return Message(message_pieces=[piece]) diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 9353989b8f..605e804ad2 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -14,13 +14,10 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, - ContentScorable, Identifiable, Message, MessagePiece, - MessageScorable, PromptResponseError, - Scorable, Score, ScorerEvaluationIdentifier, ScorerIdentifier, @@ -29,6 +26,7 @@ ) from pyrit.prompt_target.batch_helper import batch_task_async from pyrit.prompt_target.common.target_requirements import TargetRequirements +from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable if TYPE_CHECKING: from collections.abc import Sequence @@ -331,6 +329,7 @@ def _resolve_score_inputs( return scorable, expectation + @abstractmethod async def _score_scorable_async( self, *, @@ -356,42 +355,6 @@ async def _score_scorable_async( Raises: TypeError: If the scorer does not support this kind of scorable. """ - raise TypeError(f"{self.__class__.__name__} does not support scorable {type(scorable).__name__}.") - - async def _score_async(self, message: Message, *, objective: str | None = None) -> list[Score]: - """ - Score the given request response asynchronously. - - This default implementation scores all supported pieces in the message - and returns a flattened list of scores. Subclasses can override this method - to implement custom scoring logic (e.g., aggregating scores). - - Args: - message (Message): The message to score. - objective (str | None): The objective to evaluate against. Defaults to None. - - Returns: - list[Score]: A list of Score objects. - """ - if not message.message_pieces: - return [] - - # Score only the supported pieces - supported_pieces = self._get_supported_pieces(message) - - tasks = [self._score_piece_async(message_piece=piece, objective=objective) for piece in supported_pieces] - - if not tasks: - return [] - - # Run all piece-level scorings concurrently - piece_score_lists = await asyncio.gather(*tasks) - - # Flatten list[list[Score]] -> list[Score] - return [score for sublist in piece_score_lists for score in sublist] - - @abstractmethod - async def _score_piece_async(self, message_piece: MessagePiece, *, objective: str | None = None) -> list[Score]: raise NotImplementedError @staticmethod @@ -516,17 +479,6 @@ def _apply_blocked_content_substitution(self, message: Message) -> Message: return Message(message_pieces=new_pieces) - def _get_supported_pieces(self, message: Message) -> list[MessagePiece]: - """ - Get a list of supported message pieces for this scorer. - - Returns: - list[MessagePiece]: List of message pieces that are supported by this scorer's validator. - """ - return [ - piece for piece in message.message_pieces if self._validator.is_message_piece_supported(message_piece=piece) - ] - @abstractmethod def _build_fallback_score( self, *, message: Message, objective: str | None, scorer_response_blocked: bool = False diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index d142809df1..e677502591 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -7,17 +7,10 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ( - ChatMessageRole, - ComponentIdentifier, - Message, - MessagePiece, - MessageScorable, - Score, - ScoringExpectation, -) +from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleAggregatorFunc, FloatScaleScoreAggregator from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer +from pyrit.score.scorable import MessageScorable from pyrit.score.score_utils import ORIGINAL_FLOAT_VALUE_KEY from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index a8334ad520..a3583ad783 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -7,15 +7,8 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ( - ChatMessageRole, - ComponentIdentifier, - Message, - MessagePiece, - MessageScorable, - Score, - ScoringExpectation, -) +from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.score.scorable import MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc from pyrit.score.true_false.true_false_scorer import TrueFalseScorer diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index ef2af9072f..a8975dfef7 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -7,15 +7,8 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ( - ChatMessageRole, - ComponentIdentifier, - Message, - MessagePiece, - MessageScorable, - Score, - ScoringExpectation, -) +from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.score.scorable import MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer diff --git a/tests/partner_integration/azure_ai_evaluation/test_scorer_contract.py b/tests/partner_integration/azure_ai_evaluation/test_scorer_contract.py index 2bc9f5bee3..ff0c0d83cb 100644 --- a/tests/partner_integration/azure_ai_evaluation/test_scorer_contract.py +++ b/tests/partner_integration/azure_ai_evaluation/test_scorer_contract.py @@ -8,9 +8,16 @@ - RAIServiceScorer extends TrueFalseScorer Both are critical for scoring attack results. + +Scorer is now agnostic about what it scores: it takes a scorable and requires +``_score_scorable_async``. Every message-shaped hook, ``_score_piece_async`` included, moved +to ``MessageScorer``. A scorer that implements ``_score_piece_async`` must therefore extend +``MessageScorer`` instead of ``Scorer``. ``TrueFalseScorer`` already does, so +``RAIServiceScorer`` needs no change; ``AzureRAIServiceTrueFalseScorer`` does. """ from pyrit.score import ScorerPromptValidator +from pyrit.score.message_scorer import MessageScorer from pyrit.score.scorer import Scorer from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -18,9 +25,14 @@ class TestScorerContract: """Validate Scorer base class interface stability.""" - def test_scorer_has_score_piece_async(self): - """Scorer subclasses must implement _score_piece_async.""" - assert hasattr(Scorer, "_score_piece_async") + def test_scorer_requires_score_scorable_async(self): + """Scorer subclasses must implement _score_scorable_async.""" + assert "_score_scorable_async" in Scorer.__abstractmethods__ + + def test_message_scorer_has_score_piece_async(self): + """Message-shaped scorers implement _score_piece_async and must extend MessageScorer.""" + assert hasattr(MessageScorer, "_score_piece_async") + assert not hasattr(Scorer, "_score_piece_async") def test_scorer_has_validate_return_scores(self): """Scorer subclasses must implement validate_return_scores.""" @@ -38,6 +50,10 @@ def test_true_false_scorer_extends_scorer(self): """RAIServiceScorer extends TrueFalseScorer which extends Scorer.""" assert issubclass(TrueFalseScorer, Scorer) + def test_true_false_scorer_extends_message_scorer(self): + """RAIServiceScorer keeps its _score_piece_async hook through MessageScorer.""" + assert issubclass(TrueFalseScorer, MessageScorer) + def test_true_false_scorer_has_validate_return_scores(self): """TrueFalseScorer implements validate_return_scores.""" assert hasattr(TrueFalseScorer, "validate_return_scores") diff --git a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py index b23394d92c..863a25076d 100644 --- a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py +++ b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py @@ -37,13 +37,12 @@ ConversationType, Message, MessagePiece, - MessageScorable, Score, SeedPrompt, ) from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName, PromptTarget -from pyrit.score import FloatScaleThresholdScorer, Scorer, TrueFalseScorer +from pyrit.score import FloatScaleThresholdScorer, MessageScorable, Scorer, TrueFalseScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.score_utils import normalize_score_to_float diff --git a/tests/unit/models/test_expectation.py b/tests/unit/models/test_expectation.py new file mode 100644 index 0000000000..8635d0622a --- /dev/null +++ b/tests/unit/models/test_expectation.py @@ -0,0 +1,23 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import dataclasses + +import pytest + +from pyrit.models import ScoringExpectation + + +def test_expectation_defaults(): + assert ScoringExpectation().objective is None + + +def test_expectation_is_frozen(): + expectation = ScoringExpectation(objective="exfiltrate") + + with pytest.raises(dataclasses.FrozenInstanceError): + expectation.objective = "something else" + + +def test_expectations_with_equal_values_compare_equal(): + assert ScoringExpectation(objective="a") == ScoringExpectation(objective="a") diff --git a/tests/unit/models/test_scorable.py b/tests/unit/models/test_scorable.py deleted file mode 100644 index aa858a1a8d..0000000000 --- a/tests/unit/models/test_scorable.py +++ /dev/null @@ -1,63 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -import dataclasses -import uuid - -import pytest - -from pyrit.models import ContentScorable, MessagePiece, MessageReferenceScorable, MessageScorable, ScoringExpectation - - -def _message(): - return MessagePiece(role="assistant", original_value="hello").to_message() - - -@pytest.mark.parametrize( - "scorable, field_name", - [ - (MessageScorable(message=_message()), "message"), - (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), - (ContentScorable(value="hello"), "value"), - ], -) -def test_scorable_is_frozen(scorable, field_name): - with pytest.raises(dataclasses.FrozenInstanceError): - setattr(scorable, field_name, "changed") - - -def test_message_scorable_carries_the_message(): - message = _message() - scorable = MessageScorable(message=message, role_filter="assistant", skip_on_error_result=True) - - assert scorable.message is message - assert scorable.role_filter == "assistant" - assert scorable.skip_on_error_result is True - - -def test_message_reference_scorable_defaults(): - piece_id = uuid.uuid4() - scorable = MessageReferenceScorable(message_piece_ids=(piece_id,)) - - assert scorable.message_piece_ids == (piece_id,) - assert scorable.role_filter is None - assert scorable.skip_on_error_result is False - - -def test_content_scorable_defaults_to_text(): - assert ContentScorable(value="hello").data_type == "text" - - -def test_expectation_defaults(): - assert ScoringExpectation().objective is None - - -def test_expectation_is_frozen(): - expectation = ScoringExpectation(objective="exfiltrate") - - with pytest.raises(dataclasses.FrozenInstanceError): - expectation.objective = "something else" - - -def test_expectations_with_equal_values_compare_equal(): - assert ScoringExpectation(objective="a") == ScoringExpectation(objective="a") diff --git a/tests/unit/score/test_azure_content_filter.py b/tests/unit/score/test_azure_content_filter.py index da59ed00f2..48a110c10d 100644 --- a/tests/unit/score/test_azure_content_filter.py +++ b/tests/unit/score/test_azure_content_filter.py @@ -12,7 +12,8 @@ from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece, MessageScorable +from pyrit.models import Message, MessagePiece +from pyrit.score import MessageScorable from pyrit.score.float_scale.azure_content_filter_scorer import AzureContentFilterScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer diff --git a/tests/unit/score/test_conversation_history_scorer.py b/tests/unit/score/test_conversation_history_scorer.py index fb4be72386..4671196a6c 100644 --- a/tests/unit/score/test_conversation_history_scorer.py +++ b/tests/unit/score/test_conversation_history_scorer.py @@ -7,8 +7,14 @@ import pytest from pyrit.memory import CentralMemory -from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score -from pyrit.score import Scorer, SelfAskGeneralFloatScaleScorer, create_conversation_scorer +from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score +from pyrit.score import ( + MessageScorable, + MessageScorer, + Scorer, + SelfAskGeneralFloatScaleScorer, + create_conversation_scorer, +) from pyrit.score.conversation_scorer import ConversationScorer from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -49,8 +55,8 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st return [] -class MockUnsupportedScorer(Scorer): - """Mock unsupported Scorer for testing error cases""" +class MockUnsupportedScorer(MessageScorer): + """Mock scorer that is neither a FloatScaleScorer nor a TrueFalseScorer""" def __init__(self): super().__init__(validator=ScorerPromptValidator(supported_data_types=["text"])) diff --git a/tests/unit/score/test_float_scale_threshold_scorer.py b/tests/unit/score/test_float_scale_threshold_scorer.py index 7e47e30c14..69500aa662 100644 --- a/tests/unit/score/test_float_scale_threshold_scorer.py +++ b/tests/unit/score/test_float_scale_threshold_scorer.py @@ -7,8 +7,8 @@ import pytest from pyrit.memory import CentralMemory, MemoryInterface -from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score -from pyrit.score import FloatScaleThresholdScorer +from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score +from pyrit.score import FloatScaleThresholdScorer, MessageScorable from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.scorer_prompt_validator import ScorerPromptValidator diff --git a/tests/unit/score/test_gandalf_scorer.py b/tests/unit/score/test_gandalf_scorer.py index 45efd3da7a..487ce087dd 100644 --- a/tests/unit/score/test_gandalf_scorer.py +++ b/tests/unit/score/test_gandalf_scorer.py @@ -9,9 +9,9 @@ from pyrit.exceptions.exception_classes import PyritException from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece, MessageScorable +from pyrit.models import Message, MessagePiece from pyrit.prompt_target import GandalfLevel -from pyrit.score import GandalfScorer +from pyrit.score import GandalfScorer, MessageScorable def generate_password_extraction_response(response_text: str, conversation_id: str | None = None) -> Message: diff --git a/tests/unit/score/test_insecure_code_scorer.py b/tests/unit/score/test_insecure_code_scorer.py index 8404e7f746..7120148841 100644 --- a/tests/unit/score/test_insecure_code_scorer.py +++ b/tests/unit/score/test_insecure_code_scorer.py @@ -6,17 +6,9 @@ import pytest from pyrit.exceptions.exception_classes import InvalidJsonException -from pyrit.models import ( - ComponentIdentifier, - Message, - MessagePiece, - MessageScorable, - Score, - SeedPrompt, - UnvalidatedScore, -) +from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, SeedPrompt, UnvalidatedScore from pyrit.prompt_target import PromptTarget -from pyrit.score import InsecureCodeScorer +from pyrit.score import InsecureCodeScorer, MessageScorable @pytest.fixture diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py index 93d532201f..e12f9f9f47 100644 --- a/tests/unit/score/test_message_scorer.py +++ b/tests/unit/score/test_message_scorer.py @@ -7,22 +7,22 @@ import pytest from pyrit.memory import MemoryInterface -from pyrit.models import ( - ComponentIdentifier, +from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.score import ( ContentScorable, - Message, - MessagePiece, MessageReferenceScorable, MessageScorable, - Score, - ScoringExpectation, + MessageScorer, + Scorable, + Scorer, + ScorerPromptValidator, + TrueFalseScorer, ) -from pyrit.score import ScorerPromptValidator, TrueFalseScorer from pyrit.score.message_scorer import extract_objective_from_previous_turn @dataclasses.dataclass(frozen=True) -class UnsupportedScorable: +class UnsupportedScorable(Scorable): """A scorable kind no message scorer handles.""" uri: str @@ -108,6 +108,19 @@ async def test_message_reference_scorable_not_in_memory_raises(self): with pytest.raises(ValueError, match="No message pieces found in memory"): await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(missing_id,))) + async def test_message_reference_scorable_partially_in_memory_raises(self, sqlite_instance: MemoryInterface): + """A partial resolution is a caller error, so it must not be scored silently.""" + stored = _assistant_message("stored response") + sqlite_instance.add_message_to_memory(request=stored) + stored_id = stored.get_piece().id + missing_id = uuid.uuid4() + scorer = RecordingScorer() + + with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): + await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(stored_id, missing_id))) + + assert scorer.scored_messages == [] + async def test_message_reference_scorable_spanning_messages_raises(self, sqlite_instance: MemoryInterface): conversation_id = str(uuid.uuid4()) first = MessagePiece( @@ -146,6 +159,24 @@ async def test_unsupported_scorable_raises_type_error(self): await scorer.score_async(scorable=UnsupportedScorable(uri="/tmp/out.txt")) # type: ignore[arg-type] +class TestScorerBaseIsScorableAgnostic: + """Scorer knows nothing about messages; the message hooks belong to MessageScorer.""" + + def test_scorer_requires_a_scorable_implementation(self): + # A scorer that implements only the message hooks cannot be instantiated. Without + # this, such a scorer builds fine and fails later with a confusing TypeError. + assert "_score_scorable_async" in Scorer.__abstractmethods__ + + @pytest.mark.parametrize("hook", ["_score_async", "_score_piece_async", "_get_supported_pieces"]) + def test_message_hooks_live_on_message_scorer(self, hook): + assert not hasattr(Scorer, hook) + assert hasattr(MessageScorer, hook) + + def test_message_scorer_satisfies_the_scorable_contract(self): + assert "_score_scorable_async" not in MessageScorer.__abstractmethods__ + assert "_score_piece_async" in MessageScorer.__abstractmethods__ + + @pytest.mark.usefixtures("patch_central_database") class TestScorableFilters: """role_filter and skip_on_error_result are fields on the scorable, not call parameters.""" diff --git a/tests/unit/score/test_plagiarism_scorer.py b/tests/unit/score/test_plagiarism_scorer.py index ecf612ec8d..eed3048b50 100644 --- a/tests/unit/score/test_plagiarism_scorer.py +++ b/tests/unit/score/test_plagiarism_scorer.py @@ -7,8 +7,8 @@ from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece, MessageScorable -from pyrit.score import PlagiarismMetric, PlagiarismScorer +from pyrit.models import MessagePiece +from pyrit.score import MessageScorable, PlagiarismMetric, PlagiarismScorer @pytest.mark.usefixtures("patch_central_database") diff --git a/tests/unit/score/test_question_answer_scorer.py b/tests/unit/score/test_question_answer_scorer.py index 5113f73da3..2654c6b452 100644 --- a/tests/unit/score/test_question_answer_scorer.py +++ b/tests/unit/score/test_question_answer_scorer.py @@ -10,8 +10,8 @@ from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece, MessageScorable -from pyrit.score import QuestionAnswerScorer +from pyrit.models import Message, MessagePiece +from pyrit.score import MessageScorable, QuestionAnswerScorer @pytest.fixture diff --git a/tests/unit/score/test_scorable.py b/tests/unit/score/test_scorable.py new file mode 100644 index 0000000000..3aed9c7139 --- /dev/null +++ b/tests/unit/score/test_scorable.py @@ -0,0 +1,160 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import dataclasses +import uuid +from unittest.mock import MagicMock + +import pytest + +from pyrit.memory import MemoryInterface +from pyrit.models import MessagePiece +from pyrit.score import ContentScorable, MessageReferenceScorable, MessageScorable, Scorable, SingleMessageScorable + + +def _message(): + return MessagePiece(role="assistant", original_value="hello").to_message() + + +def _stored_message(value: str = "stored response"): + """Return a message that memory will accept, so it needs a conversation id.""" + return MessagePiece( + role="assistant", + original_value=value, + conversation_id=str(uuid.uuid4()), + ).to_message() + + +def _no_memory() -> MemoryInterface: + """Return a memory that the test can assert was never touched.""" + return MagicMock(spec=MemoryInterface) + + +@pytest.mark.parametrize( + "scorable, field_name", + [ + (MessageScorable(message=_message()), "message"), + (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), + (ContentScorable(value="hello"), "value"), + ], +) +def test_scorable_is_frozen(scorable, field_name): + with pytest.raises(dataclasses.FrozenInstanceError): + setattr(scorable, field_name, "changed") + + +@pytest.mark.parametrize( + "scorable", + [ + MessageScorable(message=_message()), + MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), + ContentScorable(value="hello"), + ], +) +def test_message_scorables_share_one_base(scorable): + assert isinstance(scorable, SingleMessageScorable) + assert isinstance(scorable, Scorable) + + +def test_single_message_scorable_cannot_be_instantiated(): + # The base only declares the contract; a scorable must say how it resolves. + assert "resolve_message" in SingleMessageScorable.__abstractmethods__ + + with pytest.raises(TypeError): + SingleMessageScorable() # type: ignore[abstract] + + +def test_scorables_are_keyword_only(): + with pytest.raises(TypeError): + ContentScorable("hello") # type: ignore[misc] + + +class TestMessageScorable: + def test_carries_the_message_and_filters(self): + message = _message() + scorable = MessageScorable(message=message, role_filter="assistant", skip_on_error_result=True) + + assert scorable.message is message + assert scorable.role_filter == "assistant" + assert scorable.skip_on_error_result is True + + def test_resolves_to_the_message_it_holds_without_memory(self): + message = _message() + memory = _no_memory() + + assert MessageScorable(message=message).resolve_message(memory=memory) is message + memory.get_message_pieces.assert_not_called() # type: ignore[attr-defined] + + +class TestMessageReferenceScorable: + def test_defaults(self): + piece_id = uuid.uuid4() + scorable = MessageReferenceScorable(message_piece_ids=(piece_id,)) + + assert scorable.message_piece_ids == (piece_id,) + assert scorable.role_filter is None + assert scorable.skip_on_error_result is False + + def test_resolves_from_memory(self, sqlite_instance: MemoryInterface): + stored = _stored_message() + sqlite_instance.add_message_to_memory(request=stored) + scorable = MessageReferenceScorable(message_piece_ids=(stored.get_piece().id,)) + + assert scorable.resolve_message(memory=sqlite_instance).get_value() == "stored response" + + def test_raises_when_nothing_is_in_memory(self, sqlite_instance: MemoryInterface): + missing_id = uuid.uuid4() + scorable = MessageReferenceScorable(message_piece_ids=(missing_id,)) + + with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): + scorable.resolve_message(memory=sqlite_instance) + + def test_raises_naming_only_the_missing_ids(self, sqlite_instance: MemoryInterface): + """A partial resolution is a caller error, and the error must point at the bad ids.""" + stored = _stored_message() + sqlite_instance.add_message_to_memory(request=stored) + stored_id = stored.get_piece().id + missing_id = uuid.uuid4() + scorable = MessageReferenceScorable(message_piece_ids=(stored_id, missing_id)) + + with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): + scorable.resolve_message(memory=sqlite_instance) + + def test_raises_when_pieces_span_more_than_one_message(self, sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + first = MessagePiece( + role="user", original_value="ask", conversation_id=conversation_id, sequence=0 + ).to_message() + second = MessagePiece( + role="assistant", original_value="answer", conversation_id=conversation_id, sequence=1 + ).to_message() + sqlite_instance.add_message_to_memory(request=first) + sqlite_instance.add_message_to_memory(request=second) + scorable = MessageReferenceScorable( + message_piece_ids=(first.get_piece().id, second.get_piece().id), + ) + + with pytest.raises(ValueError, match="exactly one message"): + scorable.resolve_message(memory=sqlite_instance) + + +class TestContentScorable: + def test_defaults_to_text(self): + assert ContentScorable(value="hello").data_type == "text" + + def test_resolves_to_an_unpersisted_message_without_memory(self): + """Loose scoring must stay usable with no memory behind it.""" + memory = _no_memory() + + message = ContentScorable(value="loose text").resolve_message(memory=memory) + + piece = message.get_piece() + assert piece.original_value == "loose text" + assert piece.role == "user" + assert piece.not_in_memory is True + memory.get_message_pieces.assert_not_called() # type: ignore[attr-defined] + + def test_resolves_non_text_data_types(self): + message = ContentScorable(value="path/to.png", data_type="image_path").resolve_message(memory=_no_memory()) + + assert message.get_piece().original_value_data_type == "image_path" diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index ad9c4e70f7..d453fdd54b 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -11,11 +11,12 @@ from pyrit.exceptions import InvalidJsonException, remove_markdown_json from pyrit.memory import MemoryInterface -from pyrit.models import ComponentIdentifier, Message, MessagePiece, MessageScorable, Score, ScoringExpectation +from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.prompt_target import PromptTarget from pyrit.score import ( FloatScaleScorer, JsonSchemaResponseHandler, + MessageScorable, MessageScorer, Scorer, ScorerPromptValidator, diff --git a/tests/unit/score/test_self_ask_category.py b/tests/unit/score/test_self_ask_category.py index 5da94f78f7..cbbc9e962c 100644 --- a/tests/unit/score/test_self_ask_category.py +++ b/tests/unit/score/test_self_ask_category.py @@ -11,8 +11,14 @@ from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece, MessageScorable -from pyrit.score import ContentClassifier, ContentClassifierCategory, ContentClassifierPaths, SelfAskCategoryScorer +from pyrit.models import Message, MessagePiece +from pyrit.score import ( + ContentClassifier, + ContentClassifierCategory, + ContentClassifierPaths, + MessageScorable, + SelfAskCategoryScorer, +) HARM_CLASSIFIER = ContentClassifier.from_yaml(ContentClassifierPaths.HARMFUL_CONTENT_CLASSIFIER.value) diff --git a/tests/unit/score/test_self_ask_question_answer_scorer.py b/tests/unit/score/test_self_ask_question_answer_scorer.py index a849338116..09024ef381 100644 --- a/tests/unit/score/test_self_ask_question_answer_scorer.py +++ b/tests/unit/score/test_self_ask_question_answer_scorer.py @@ -5,8 +5,9 @@ import pytest -from pyrit.models import ComponentIdentifier, MessagePiece, MessageScorable, Score, ScoringExpectation, UnvalidatedScore +from pyrit.models import ComponentIdentifier, MessagePiece, Score, ScoringExpectation, UnvalidatedScore from pyrit.prompt_target import PromptTarget +from pyrit.score import MessageScorable from pyrit.score.true_false.self_ask_question_answer_scorer import SelfAskQuestionAnswerScorer diff --git a/tests/unit/score/test_self_ask_refusal.py b/tests/unit/score/test_self_ask_refusal.py index 217a025ac5..50a43bd2bf 100644 --- a/tests/unit/score/test_self_ask_refusal.py +++ b/tests/unit/score/test_self_ask_refusal.py @@ -19,10 +19,9 @@ JsonResponseConfig, Message, MessagePiece, - MessageScorable, SeedPrompt, ) -from pyrit.score import JsonSchemaResponseHandler, RefusalScorerPaths, SelfAskRefusalScorer +from pyrit.score import JsonSchemaResponseHandler, MessageScorable, RefusalScorerPaths, SelfAskRefusalScorer @pytest.fixture diff --git a/tests/unit/score/test_shieldgemma_scorer.py b/tests/unit/score/test_shieldgemma_scorer.py index f365cc7ee3..6c1aef2438 100644 --- a/tests/unit/score/test_shieldgemma_scorer.py +++ b/tests/unit/score/test_shieldgemma_scorer.py @@ -9,9 +9,10 @@ from pyrit.exceptions import InvalidJsonException from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import JSON_SCHEMA_METADATA_KEY, Message, MessagePiece, MessageScorable +from pyrit.models import JSON_SCHEMA_METADATA_KEY, Message, MessagePiece from pyrit.prompt_target import PromptTarget from pyrit.score import ( + MessageScorable, ShieldGemmaGuideline, ShieldGemmaMessageRole, ShieldGemmaPolicy, diff --git a/tests/unit/score/test_substring.py b/tests/unit/score/test_substring.py index ff014eead8..6f65f024aa 100644 --- a/tests/unit/score/test_substring.py +++ b/tests/unit/score/test_substring.py @@ -10,8 +10,8 @@ from pyrit.analytics import ApproximateTextMatching, ExactTextMatching from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece, MessageScorable -from pyrit.score import SubStringScorer +from pyrit.models import MessagePiece +from pyrit.score import MessageScorable, SubStringScorer @pytest.fixture diff --git a/tests/unit/score/test_true_false_composite_scorer.py b/tests/unit/score/test_true_false_composite_scorer.py index e3b73f1599..3656fe5a9c 100644 --- a/tests/unit/score/test_true_false_composite_scorer.py +++ b/tests/unit/score/test_true_false_composite_scorer.py @@ -6,8 +6,14 @@ import pytest from pyrit.memory.central_memory import CentralMemory -from pyrit.models import ComponentIdentifier, MessagePiece, MessageScorable, Score, ScoringExpectation -from pyrit.score import FloatScaleScorer, TrueFalseCompositeScorer, TrueFalseScoreAggregator, TrueFalseScorer +from pyrit.models import ComponentIdentifier, MessagePiece, Score, ScoringExpectation +from pyrit.score import ( + FloatScaleScorer, + MessageScorable, + TrueFalseCompositeScorer, + TrueFalseScoreAggregator, + TrueFalseScorer, +) def _mock_scorer_id(name: str = "MockScorer") -> ComponentIdentifier: diff --git a/tests/unit/score/test_true_false_inverter.py b/tests/unit/score/test_true_false_inverter.py index 7bd459f3da..ecb0bd924d 100644 --- a/tests/unit/score/test_true_false_inverter.py +++ b/tests/unit/score/test_true_false_inverter.py @@ -9,8 +9,8 @@ from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import MessagePiece, MessageScorable -from pyrit.score import SubStringScorer, TrueFalseInverterScorer +from pyrit.models import MessagePiece +from pyrit.score import MessageScorable, SubStringScorer, TrueFalseInverterScorer @pytest.fixture From 27ddea09958dbb93b57d6b1bc81c48268279088a Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 10:41:44 -0700 Subject: [PATCH 04/11] Stop treating loose content as a scorable that names a message ContentScorable was made a SingleMessageScorable, which claimed it names exactly one Message. It does not. It carries content, and the Message it produced was fabricated on the spot as an adapter so message scorers could read it. The design proposal is explicit that a scorable is a reference and that loose content is the one exception (section 4), and that a later phase persists the content as a row of its own and retires the fabricated, never-persisted message (section 9.1). Putting it under that base also handed it role_filter and skip_on_error_result, which are meaningless for content that has no role and no error state. That was not only untidy: ContentScorable(value=..., role_filter="assistant") was accepted and then silently scored nothing, because the adapted piece is always role="user". The previous getattr duck-typing defaulted the filter to None, so this state was newly reachable. ContentScorable now derives from Scorable directly and exposes to_ephemeral_message instead of resolve_message. The different verb is the point: the message-shaped scorables return a message that already exists, while loose content becomes one. MessageScorer bridges the two shapes in two arms, and the content arm is marked as the transitional adapter it is. Passing a filter to ContentScorable is now a TypeError at construction. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/score/message_scorer.py | 19 ++++++++++----- pyrit/score/scorable.py | 39 ++++++++++++++++++++----------- tests/unit/score/test_scorable.py | 34 ++++++++++++++++++++------- 3 files changed, 63 insertions(+), 29 deletions(-) diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index fe8124b19e..c1715d2f30 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException -from pyrit.score.scorable import Scorable, SingleMessageScorable +from pyrit.score.scorable import ContentScorable, Scorable, SingleMessageScorable from pyrit.score.scorer import Scorer if TYPE_CHECKING: @@ -79,7 +79,7 @@ async def _score_scorable_async( Resolve a message scorable and score the message it names. Args: - scorable (Scorable): Any ``SingleMessageScorable``. + scorable (Scorable): A ``SingleMessageScorable`` or a ``ContentScorable``. expectation (ScoringExpectation | None): What to look for. infer_objective_from_request (bool): Deprecated; read the objective from the previous turn when the expectation carries none. @@ -94,15 +94,22 @@ async def _score_scorable_async( PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). """ - if not isinstance(scorable, SingleMessageScorable): + if isinstance(scorable, SingleMessageScorable): + message = scorable.resolve_message(memory=self._memory) + role_filter = scorable.role_filter + skip_on_error_result = scorable.skip_on_error_result + elif isinstance(scorable, ContentScorable): + # Loose content names no message, so there is nothing to filter on. Phase 2 + # persists the content and this arm goes away. + message = scorable.to_ephemeral_message() + role_filter = None + skip_on_error_result = False + else: raise TypeError( f"{self.__class__.__name__} scores messages, so it cannot score {type(scorable).__name__}. " "Pass a MessageScorable, a MessageReferenceScorable, or a ContentScorable." ) - message = scorable.resolve_message(memory=self._memory) - role_filter = scorable.role_filter - skip_on_error_result = scorable.skip_on_error_result objective = expectation.objective if expectation else None # Structured refusals are persisted as blocked error pieces, but scorers should diff --git a/pyrit/score/scorable.py b/pyrit/score/scorable.py index 48da9408c9..3c04853537 100644 --- a/pyrit/score/scorable.py +++ b/pyrit/score/scorable.py @@ -20,23 +20,27 @@ from pyrit.memory import MemoryInterface -class Scorable(ABC): # noqa: B024 root type; each scorer family declares its own resolution +class Scorable(ABC): # noqa: B024 root type; each scorer family declares its own contract """ What a scorer looks at. - A scorable is a reference rather than a payload: it names the evidence instead of - carrying it. Each scorer accepts the scorable kinds it supports and rejects the rest. + A scorable is normally a reference: it names the evidence instead of carrying it, and + the scorer resolves it. ``ContentScorable`` is the exception, because loose content has + nothing behind it to point at. Each scorer accepts the scorable kinds it supports and + rejects the rest. """ @dataclass(frozen=True, kw_only=True) class SingleMessageScorable(Scorable, ABC): """ - A scorable that names exactly one message. + A scorable that names exactly one message that already exists. - Subclasses resolve themselves. Only identity resolution belongs here — returning the - thing the scorable names. A scorer that derives a different kind of evidence from the - same reference, such as trace ids or a file write, does that widening itself. + The message is in hand or in memory, so it has a real role and a real error state and + the filters below mean something. Subclasses resolve themselves, but only identity + resolution belongs here — returning the message the scorable names. A scorer that + derives a different kind of evidence from the same reference, such as trace ids or a + file write, does that widening itself. """ role_filter: ChatMessageRole | None = None @@ -122,21 +126,28 @@ def resolve_message(self, *, memory: MemoryInterface) -> Message: @dataclass(frozen=True, kw_only=True) -class ContentScorable(SingleMessageScorable): - """Loose content with no conversation behind it.""" +class ContentScorable(Scorable): + """ + Loose content with no conversation behind it. + + This names content, not a message, so it is not a ``SingleMessageScorable``: there is no + role and no error state, and the filters those scorables carry would be meaningless here. + A message scorer adapts it with ``to_ephemeral_message``. + """ value: str data_type: PromptDataType = "text" - def resolve_message(self, *, memory: MemoryInterface) -> Message: + def to_ephemeral_message(self) -> Message: """ - Return a message holding the loose content, marked as never persisted. + Wrap the content as a message so a message scorer can read it. - Args: - memory (MemoryInterface): Unused; loose content has no conversation to read. + This is an adapter, not a resolution: no such message exists until this call builds + one, and it is marked as never persisted. Phase 2 stores loose content as a row of + its own and retires this. Returns: - Message: The message to score. + Message: A throwaway message holding the content. """ piece = MessagePiece( role="user", diff --git a/tests/unit/score/test_scorable.py b/tests/unit/score/test_scorable.py index 3aed9c7139..c3ab4c8c71 100644 --- a/tests/unit/score/test_scorable.py +++ b/tests/unit/score/test_scorable.py @@ -48,14 +48,29 @@ def test_scorable_is_frozen(scorable, field_name): [ MessageScorable(message=_message()), MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), - ContentScorable(value="hello"), ], ) -def test_message_scorables_share_one_base(scorable): +def test_scorables_naming_a_message_share_one_base(scorable): assert isinstance(scorable, SingleMessageScorable) assert isinstance(scorable, Scorable) +def test_content_scorable_names_content_rather_than_a_message(): + """Loose content has no message behind it, so it must not inherit message identity.""" + scorable = ContentScorable(value="hello") + + assert isinstance(scorable, Scorable) + assert not isinstance(scorable, SingleMessageScorable) + + +@pytest.mark.parametrize("field_name", ["role_filter", "skip_on_error_result"]) +def test_content_scorable_rejects_message_filters(field_name): + # Loose content has no role and no error state. Accepting a role filter here used to + # score nothing at all, silently, because the adapted message is always role="user". + with pytest.raises(TypeError): + ContentScorable(value="hello", **{field_name: "assistant"}) + + def test_single_message_scorable_cannot_be_instantiated(): # The base only declares the contract; a scorable must say how it resolves. assert "resolve_message" in SingleMessageScorable.__abstractmethods__ @@ -142,19 +157,20 @@ class TestContentScorable: def test_defaults_to_text(self): assert ContentScorable(value="hello").data_type == "text" - def test_resolves_to_an_unpersisted_message_without_memory(self): + def test_adapts_to_an_unpersisted_message(self): """Loose scoring must stay usable with no memory behind it.""" - memory = _no_memory() - - message = ContentScorable(value="loose text").resolve_message(memory=memory) + message = ContentScorable(value="loose text").to_ephemeral_message() piece = message.get_piece() assert piece.original_value == "loose text" assert piece.role == "user" assert piece.not_in_memory is True - memory.get_message_pieces.assert_not_called() # type: ignore[attr-defined] - def test_resolves_non_text_data_types(self): - message = ContentScorable(value="path/to.png", data_type="image_path").resolve_message(memory=_no_memory()) + def test_adapts_non_text_data_types(self): + message = ContentScorable(value="path/to.png", data_type="image_path").to_ephemeral_message() assert message.get_piece().original_value_data_type == "image_path" + + def test_does_not_resolve_a_message(self): + # It has no message to name, so it must not answer the resolution contract. + assert not hasattr(ContentScorable(value="hello"), "resolve_message") From 1314a62718a4892aaaad64381692a7e1a6e2448a Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 13:09:47 -0700 Subject: [PATCH 05/11] Address review notes on the scorer contract Collapse the two message scorables into one by-id MessageScorable, per the gist. Use the from_message classmethods instead of a free helper. Correct the legacy removal version to 2.0.0. Rename _resolve_score_inputs to _consolidate_legacy_inputs. Deprecate extract_objective_from_previous_turn. Remove the dead objective ternary in Crescendo. Fail when a scorable names a piece id that memory does not hold. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/executor/attack/multi_turn/crescendo.py | 4 +- .../executor/attack/multi_turn/red_teaming.py | 2 +- pyrit/score/__init__.py | 10 +- pyrit/score/audio_transcript_scorer.py | 2 +- pyrit/score/message_scorer.py | 16 +- pyrit/score/scorable.py | 132 ++++++++------- pyrit/score/scorer.py | 43 +++-- .../float_scale_threshold_scorer.py | 12 +- .../true_false/true_false_composite_scorer.py | 12 +- .../true_false/true_false_inverter_scorer.py | 12 +- .../multi_turn/test_crescendo_resilience.py | 6 +- .../attack/multi_turn/test_red_teaming.py | 3 + .../attack/multi_turn/test_tree_of_attacks.py | 3 +- tests/unit/mocks.py | 57 ++++++- tests/unit/score/test_azure_content_filter.py | 8 +- .../score/test_conversation_history_scorer.py | 59 +++---- .../test_float_scale_threshold_scorer.py | 3 +- tests/unit/score/test_gandalf_scorer.py | 15 +- tests/unit/score/test_insecure_code_scorer.py | 7 +- tests/unit/score/test_message_scorer.py | 95 +++++------ tests/unit/score/test_plagiarism_scorer.py | 12 +- .../unit/score/test_question_answer_scorer.py | 47 +++--- tests/unit/score/test_scorable.py | 113 +++++++------ tests/unit/score/test_scorer.py | 158 +++++++++++------- tests/unit/score/test_self_ask_category.py | 22 ++- .../test_self_ask_question_answer_scorer.py | 4 +- tests/unit/score/test_self_ask_refusal.py | 4 +- tests/unit/score/test_shieldgemma_scorer.py | 6 +- tests/unit/score/test_substring.py | 4 +- .../score/test_true_false_composite_scorer.py | 29 ++-- tests/unit/score/test_true_false_inverter.py | 4 +- 31 files changed, 521 insertions(+), 383 deletions(-) diff --git a/pyrit/executor/attack/multi_turn/crescendo.py b/pyrit/executor/attack/multi_turn/crescendo.py index 9fb48abfec..05fad46acd 100644 --- a/pyrit/executor/attack/multi_turn/crescendo.py +++ b/pyrit/executor/attack/multi_turn/crescendo.py @@ -672,8 +672,8 @@ async def _check_refusal_async(self, context: CrescendoAttackContext, objective: objective=context.objective, ): scores = await self._refusal_scorer.score_async( - scorable=MessageScorable(message=context.last_response), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + scorable=MessageScorable.from_message(context.last_response), + expectation=ScoringExpectation(objective=objective), ) return scores[0] diff --git a/pyrit/executor/attack/multi_turn/red_teaming.py b/pyrit/executor/attack/multi_turn/red_teaming.py index 6ae5a19055..7216254cdb 100644 --- a/pyrit/executor/attack/multi_turn/red_teaming.py +++ b/pyrit/executor/attack/multi_turn/red_teaming.py @@ -525,7 +525,7 @@ async def _score_response_async(self, *, context: MultiTurnAttackContext[Any]) - ): # score_async handles blocked, filtered, other errors scoring_results = await self._objective_scorer.score_async( - scorable=MessageScorable(message=context.last_response, role_filter="assistant"), + scorable=MessageScorable.from_message(context.last_response, role_filter="assistant"), expectation=ScoringExpectation(objective=context.objective), ) diff --git a/pyrit/score/__init__.py b/pyrit/score/__init__.py index 1cabd5b035..79460fb5b0 100644 --- a/pyrit/score/__init__.py +++ b/pyrit/score/__init__.py @@ -33,13 +33,7 @@ from pyrit.score.float_scale.self_ask_scale_scorer import SelfAskScaleScorer, render_scale_system_prompt from pyrit.score.message_scorer import MessageScorer from pyrit.score.response_handler import CallableResponseHandler, JsonSchemaResponseHandler, ResponseHandler -from pyrit.score.scorable import ( - ContentScorable, - MessageReferenceScorable, - MessageScorable, - Scorable, - SingleMessageScorable, -) +from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable from pyrit.score.scorer import Scorer from pyrit.score.scorer_evaluation.metrics_type import MetricsType, RegistryUpdateBehavior from pyrit.score.scorer_evaluation.scorer_metrics import ( @@ -196,7 +190,6 @@ def __getattr__(name: str) -> object: "LlamaGuardPolicy", "LlamaGuardScorer", "MarkdownInjectionScorer", - "MessageReferenceScorable", "MessageScorable", "MessageScorer", "MethKeywordScorer", @@ -245,7 +238,6 @@ def __getattr__(name: str) -> object: "SelfAskRefusalScorer", "SelfAskScaleScorer", "SelfAskTrueFalseScorer", - "SingleMessageScorable", "ScorerPrinter", "SHIELDGEMMA_DEFAULT_POLICY_PATH", "ShieldGemmaGuideline", diff --git a/pyrit/score/audio_transcript_scorer.py b/pyrit/score/audio_transcript_scorer.py index 7ed29d505e..aaf82af688 100644 --- a/pyrit/score/audio_transcript_scorer.py +++ b/pyrit/score/audio_transcript_scorer.py @@ -187,7 +187,7 @@ async def _score_audio_async(self, *, message_piece: MessagePiece, objective: st # Score the transcript transcript_scores = await self.text_scorer.score_async( - scorable=MessageScorable(message=text_message), + scorable=MessageScorable.from_message(text_message), expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index c1715d2f30..5c6dd4dc86 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException -from pyrit.score.scorable import ContentScorable, Scorable, SingleMessageScorable +from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable from pyrit.score.scorer import Scorer if TYPE_CHECKING: @@ -23,8 +23,12 @@ def extract_objective_from_previous_turn(*, message: Message, memory: MemoryInte """ Read the text of the turn before an assistant message and use it as the objective. - This is what to look for, so it belongs to the caller that builds the expectation, not - to the scorer. It lives here because it is message-shaped. + .. deprecated:: + This conflates scoring with building an expectation. What to look for belongs to + the caller that builds the ``ScoringExpectation``, not to the scorer. It exists only + to support the deprecated ``infer_objective_from_request`` parameter, and both are + removed in the next major release. Resolve the objective at the call site and pass + it on the expectation instead. Args: message (Message): The assistant message whose previous turn supplies the objective. @@ -79,7 +83,7 @@ async def _score_scorable_async( Resolve a message scorable and score the message it names. Args: - scorable (Scorable): A ``SingleMessageScorable`` or a ``ContentScorable``. + scorable (Scorable): A ``MessageScorable`` or a ``ContentScorable``. expectation (ScoringExpectation | None): What to look for. infer_objective_from_request (bool): Deprecated; read the objective from the previous turn when the expectation carries none. @@ -94,7 +98,7 @@ async def _score_scorable_async( PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). """ - if isinstance(scorable, SingleMessageScorable): + if isinstance(scorable, MessageScorable): message = scorable.resolve_message(memory=self._memory) role_filter = scorable.role_filter skip_on_error_result = scorable.skip_on_error_result @@ -107,7 +111,7 @@ async def _score_scorable_async( else: raise TypeError( f"{self.__class__.__name__} scores messages, so it cannot score {type(scorable).__name__}. " - "Pass a MessageScorable, a MessageReferenceScorable, or a ContentScorable." + "Pass a MessageScorable or a ContentScorable." ) objective = expectation.objective if expectation else None diff --git a/pyrit/score/scorable.py b/pyrit/score/scorable.py index 3c04853537..f47fd666b7 100644 --- a/pyrit/score/scorable.py +++ b/pyrit/score/scorable.py @@ -4,7 +4,7 @@ from __future__ import annotations import uuid # noqa: TC003 (runtime-required by dataclass field annotations) -from abc import ABC, abstractmethod +from abc import ABC from dataclasses import dataclass from typing import TYPE_CHECKING @@ -32,69 +32,24 @@ class Scorable(ABC): # noqa: B024 root type; each scorer family declares its o @dataclass(frozen=True, kw_only=True) -class SingleMessageScorable(Scorable, ABC): - """ - A scorable that names exactly one message that already exists. - - The message is in hand or in memory, so it has a real role and a real error state and - the filters below mean something. Subclasses resolve themselves, but only identity - resolution belongs here — returning the message the scorable names. A scorer that - derives a different kind of evidence from the same reference, such as trace ids or a - file write, does that widening itself. - """ - - role_filter: ChatMessageRole | None = None - skip_on_error_result: bool = False - - @abstractmethod - def resolve_message(self, *, memory: MemoryInterface) -> Message: - """ - Return the message this scorable names. - - Args: - memory (MemoryInterface): Memory to resolve references against. - - Returns: - Message: The message to score. - """ - - -@dataclass(frozen=True, kw_only=True) -class MessageScorable(SingleMessageScorable): - """ - A message the caller already holds. - - Use this when the message is in hand — an attack scoring the response it just - received, or a caller scoring a message whose pieces were never persisted. Use - ``MessageReferenceScorable`` when only the piece ids are known. - """ - - message: Message - - def resolve_message(self, *, memory: MemoryInterface) -> Message: - """ - Return the message the caller supplied. - - Args: - memory (MemoryInterface): Unused; the message is already in hand. - - Returns: - Message: The message to score. - """ - return self.message - - -@dataclass(frozen=True, kw_only=True) -class MessageReferenceScorable(SingleMessageScorable): +class MessageScorable(Scorable): """ Specific message pieces, resolved from memory by id. - Use this when the pieces are persisted and only their ids are known. Use - ``MessageScorable`` when the caller already holds the message. + This names one message, or a subset of its pieces. Loose content that was never + persisted has no ids to name, so it is a ``ContentScorable`` instead. """ message_piece_ids: tuple[uuid.UUID | str, ...] + # The design proposal puts role_filter on ConversationScorable only and leaves this + # open as its phase-1 decision, on the grounds that naming pieces is already a + # selection and filtering them again is a second one. Attacks do filter a named + # response by role today, so both fields stay here until ConversationScorable arrives + # and can own the selection instead. + role_filter: ChatMessageRole | None = None + skip_on_error_result: bool = False + def resolve_message(self, *, memory: MemoryInterface) -> Message: """ Return the single message the referenced pieces form. @@ -110,6 +65,9 @@ def resolve_message(self, *, memory: MemoryInterface) -> Message: exactly one message. """ pieces = memory.get_message_pieces(prompt_ids=list(self.message_piece_ids)) + wanted = {str(piece_id) for piece_id in self.message_piece_ids} + # Only the named pieces count, whatever else the filter happened to return. + pieces = [piece for piece in pieces if str(piece.id) in wanted] found = {str(piece.id) for piece in pieces} missing = [str(piece_id) for piece_id in self.message_piece_ids if str(piece_id) not in found] if missing: @@ -122,7 +80,42 @@ def resolve_message(self, *, memory: MemoryInterface) -> Message: f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " "Reference pieces from a single message." ) - return messages[0] + + # Memory returns pieces in its own order, but the scorable names an ordered tuple, + # and a multi-piece message reads differently if its pieces are shuffled. + resolved = messages[0] + by_id = {str(piece.id): piece for piece in resolved.message_pieces} + resolved.message_pieces = [by_id[str(piece_id)] for piece_id in self.message_piece_ids] + return resolved + + @classmethod + def from_message( + cls, + message: Message, + *, + role_filter: ChatMessageRole | None = None, + skip_on_error_result: bool = False, + ) -> MessageScorable: + """ + Name the pieces of a persisted message. + + A scorable is a reference, so a message in hand is scored by naming its pieces + rather than by travelling through the call. Use ``scorable_from_message`` instead + when the message may not be persisted. + + Args: + message (Message): The message whose pieces to name. + role_filter (ChatMessageRole | None): Only score the message when it has this role. + skip_on_error_result (bool): Skip scoring when the message is an error result. + + Returns: + MessageScorable: A scorable naming the message's pieces. + """ + return cls( + message_piece_ids=tuple(piece.id for piece in message.message_pieces), + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) @dataclass(frozen=True, kw_only=True) @@ -130,9 +123,9 @@ class ContentScorable(Scorable): """ Loose content with no conversation behind it. - This names content, not a message, so it is not a ``SingleMessageScorable``: there is no - role and no error state, and the filters those scorables carry would be meaningless here. - A message scorer adapts it with ``to_ephemeral_message``. + This names content, not a message: there is no role and no error state, so the filters + a ``MessageScorable`` carries would be meaningless here. A message scorer adapts it + with ``to_ephemeral_message``. """ value: str @@ -156,3 +149,20 @@ def to_ephemeral_message(self) -> Message: ) piece.not_in_memory = True return Message(message_pieces=[piece]) + + @classmethod + def from_message(cls, message: Message) -> ContentScorable: + """ + Take the content out of a message that was never persisted. + + A message the caller built by hand has no ids to name, so it is content rather than + a reference. Only the first piece is taken, because loose content is a single value. + + Args: + message (Message): The unpersisted message whose content to take. + + Returns: + ContentScorable: A scorable holding the message's content. + """ + piece = message.get_piece() + return cls(value=piece.original_value, data_type=piece.original_value_data_type) diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 605e804ad2..800e8139a7 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -40,7 +40,7 @@ logger = logging.getLogger(__name__) #: Release in which the message-shaped ``score_async`` parameters are removed. -LEGACY_SCORE_ASYNC_REMOVED_IN = "1.3.0" +LEGACY_SCORE_ASYNC_REMOVED_IN = "2.0.0" class Scorer(Identifiable, abc.ABC): @@ -219,11 +219,10 @@ async def score_async( ``expectation`` says what to look for. Keeping them separate is what lets an attack forward a question it does not understand to a scorer that does. - This signature is experimental for one release and may change. Args: message (Message | None): Deprecated. Pass - ``scorable=MessageScorable(message=...)`` instead. + ``scorable=MessageScorable.from_message(...)`` instead. scorable (Scorable | None): What to look at. expectation (ScoringExpectation | None): What to look for. Defaults to None. objective (str | None): Deprecated. Pass @@ -243,7 +242,7 @@ async def score_async( parameters that do not apply to them. TypeError: If this scorer does not support this kind of scorable. """ - resolved_scorable, resolved_expectation = self._resolve_score_inputs( + resolved_scorable, resolved_expectation = self._consolidate_legacy_inputs( message=message, scorable=scorable, expectation=expectation, @@ -267,7 +266,7 @@ async def score_async( return scores - def _resolve_score_inputs( + def _consolidate_legacy_inputs( self, *, message: Message | None, @@ -318,11 +317,21 @@ def _resolve_score_inputs( if scorable is None: # The caller asked about this message, not about everything stored alongside it, # so the shim maps to the exact message rather than widening to its conversation. - scorable = MessageScorable( - message=cast("Message", message), - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) + # A message the caller built by hand has nothing in memory to name, so a single + # unpersisted piece becomes loose content. Both arms go away with the parameter. + legacy_message = cast("Message", message) + pieces = legacy_message.message_pieces + if len(pieces) == 1 and pieces[0].not_in_memory: + scorable = ContentScorable( + value=pieces[0].original_value, + data_type=pieces[0].original_value_data_type, + ) + else: + scorable = MessageScorable.from_message( + legacy_message, + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) if objective is not None: expectation = ScoringExpectation(objective=objective) @@ -665,7 +674,7 @@ async def score_prompts_batch_async( return [] scorables = [ - MessageScorable(message=message, role_filter=role_filter, skip_on_error_result=skip_on_error_result) + MessageScorable.from_message(message, role_filter=role_filter, skip_on_error_result=skip_on_error_result) for message in messages ] expectations = [ScoringExpectation(objective=objective) for objective in objectives] @@ -815,8 +824,8 @@ async def score_response_async( skip_on_error_result=skip_on_error_result, ) obj_task = objective_scorer.score_async( - scorable=MessageScorable( - message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result + scorable=MessageScorable.from_message( + response, role_filter=role_filter, skip_on_error_result=skip_on_error_result ), expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) @@ -825,8 +834,8 @@ async def score_response_async( result["objective_scores"] = obj_scores else: obj_scores = await objective_scorer.score_async( - scorable=MessageScorable( - message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result + scorable=MessageScorable.from_message( + response, role_filter=role_filter, skip_on_error_result=skip_on_error_result ), expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) @@ -863,7 +872,9 @@ async def score_response_multiple_scorers_async( return [] # Create all scoring tasks, note TEMPORARY fix to prevent multi-piece responses from breaking scoring logic - scorable = MessageScorable(message=response, role_filter=role_filter, skip_on_error_result=skip_on_error_result) + scorable = MessageScorable.from_message( + response, role_filter=role_filter, skip_on_error_result=skip_on_error_result + ) expectation = ScoringExpectation(objective=objective) if objective is not None else None tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in scorers] diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index e677502591..26184a64a8 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -10,7 +10,7 @@ from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleAggregatorFunc, FloatScaleScoreAggregator from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer -from pyrit.score.scorable import MessageScorable +from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.score_utils import ORIGINAL_FLOAT_VALUE_KEY from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -99,8 +99,16 @@ async def _score_async( Returns: list[Score]: A list containing a single true/false Score object based on the threshold comparison. """ + # A message the caller built by hand has no ids to name, so it is content rather + # than a reference. Both arms go away when loose content is persisted in its own right. + pieces = message.message_pieces + scorable = ( + ContentScorable.from_message(message) + if len(pieces) == 1 and pieces[0].not_in_memory + else MessageScorable.from_message(message, role_filter=role_filter) + ) scores = await self._scorer.score_async( - scorable=MessageScorable(message=message, role_filter=role_filter), + scorable=scorable, expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index a3583ad783..0c39791d59 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -8,7 +8,7 @@ from pyrit.prompt_target import PromptTarget from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation -from pyrit.score.scorable import MessageScorable +from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -98,9 +98,17 @@ async def _score_async( ValueError: If any constituent scorer does not return exactly one score. ValueError: If no scores are generated from the request response pieces. """ + # A message the caller built by hand has no ids to name, so it is content rather + # than a reference. Both arms go away when loose content is persisted in its own right. + pieces = message.message_pieces + scorable = ( + ContentScorable.from_message(message) + if len(pieces) == 1 and pieces[0].not_in_memory + else MessageScorable.from_message(message, role_filter=role_filter) + ) tasks = [ scorer.score_async( - scorable=MessageScorable(message=message, role_filter=role_filter), + scorable=scorable, expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) for scorer in self._scorers diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index a8975dfef7..bde6bb318a 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -8,7 +8,7 @@ from pyrit.prompt_target import PromptTarget from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation -from pyrit.score.scorable import MessageScorable +from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -74,8 +74,16 @@ async def _score_async( Returns: list[Score]: A list containing a single Score object with the inverted true/false value. """ + # A message the caller built by hand has no ids to name, so it is content rather + # than a reference. Both arms go away when loose content is persisted in its own right. + pieces = message.message_pieces + scorable = ( + ContentScorable.from_message(message) + if len(pieces) == 1 and pieces[0].not_in_memory + else MessageScorable.from_message(message, role_filter=role_filter) + ) scores = await self._scorer.score_async( - scorable=MessageScorable(message=message, role_filter=role_filter), + scorable=scorable, expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) inv_score = scores[0] diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py index 2d52a943c9..ba2e604990 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py @@ -19,6 +19,7 @@ CrescendoAttackContext, CrescendoAttackResult, ) +from pyrit.memory import CentralMemory from pyrit.models import AttackOutcome, ComponentIdentifier, ConversationType, Message, MessagePiece, Score from pyrit.prompt_normalizer import PromptNormalizer from pyrit.score import Scorer, TrueFalseScorer @@ -279,8 +280,11 @@ async def score_objective(**_kwargs): assert result.conversation_id == final_conversation_id assert len({attempt.conversation_id for attempt in adversarial_target.attempts}) == 1 + # A scorable names piece ids rather than carrying the message, so read them back. + memory = CentralMemory.get_memory_instance() refusal_inputs = [ - call.kwargs["scorable"].message.get_value() for call in refusal_scorer.score_async.await_args_list + call.kwargs["scorable"].resolve_message(memory=memory).get_value() + for call in refusal_scorer.score_async.await_args_list ] assert refusal_inputs == [ "response-1", diff --git a/tests/unit/executor/attack/multi_turn/test_red_teaming.py b/tests/unit/executor/attack/multi_turn/test_red_teaming.py index 4603152590..74930594f4 100644 --- a/tests/unit/executor/attack/multi_turn/test_red_teaming.py +++ b/tests/unit/executor/attack/multi_turn/test_red_teaming.py @@ -968,9 +968,12 @@ async def test_score_response_returns_none_for_blocked( response_piece = MagicMock(spec=MessagePiece) response_piece.is_blocked.return_value = True + response_piece.id = uuid.uuid4() basic_context.last_response = MagicMock(spec=Message) basic_context.last_response.get_piece.return_value = response_piece + # A scorable names piece ids, so the mock has to expose its pieces. + basic_context.last_response.message_pieces = [response_piece] # Configure the mock scorer to return empty list for blocked response mock_objective_scorer.score_async = AsyncMock(return_value=[]) diff --git a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py index 863a25076d..4fc630bb99 100644 --- a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py +++ b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py @@ -12,6 +12,7 @@ import pytest from treelib.tree import Tree +from unit.mocks import store_message from pyrit.exceptions import InvalidJsonException from pyrit.executor.attack import ( @@ -391,7 +392,7 @@ async def create_threshold_score_async(*, original_float_value: float, threshold ) # Score using the actual FloatScaleThresholdScorer - scores = await threshold_scorer.score_async(scorable=MessageScorable(message=dummy_message)) + scores = await threshold_scorer.score_async(scorable=MessageScorable.from_message(store_message(dummy_message))) return scores[0] @staticmethod diff --git a/tests/unit/mocks.py b/tests/unit/mocks.py index 89d7cd2cd0..5a6051979d 100644 --- a/tests/unit/mocks.py +++ b/tests/unit/mocks.py @@ -10,7 +10,7 @@ from typing import Any from unittest.mock import MagicMock, patch -from pyrit.memory import AzureSQLMemory, CentralMemory, PromptMemoryEntry +from pyrit.memory import AzureSQLMemory, CentralMemory, MemoryInterface, PromptMemoryEntry from pyrit.models import ( ComponentIdentifier, Message, @@ -318,6 +318,61 @@ def get_audio_message_piece() -> MessagePiece: ) +def mock_memory_resolving(*messages: Message) -> MagicMock: + """ + Return a fake memory that can resolve the given messages by piece id. + + Scoring resolves a scorable through memory, so a fake that answers nothing makes every + score fail. This keeps the test free of a real database while still supporting lookups. + + Args: + *messages (Message): Messages the fake should be able to resolve. + + Returns: + MagicMock: A MemoryInterface stand-in with a working get_message_pieces. + """ + known = {str(piece.id): piece for message in messages for piece in message.message_pieces} + memory = MagicMock(MemoryInterface) + memory.get_message_pieces.side_effect = lambda **kwargs: [ + known[str(piece_id)] for piece_id in kwargs.get("prompt_ids", []) or [] if str(piece_id) in known + ] + return memory + + +def store_message(message: Message) -> Message: + """ + Persist a message so a scorable can name its pieces, and return it. + + A ``MessageScorable`` is a reference: it names piece ids and resolves them from memory. + Tests that build a message by hand have to store it first, or resolution has nothing to + find. This fills in whatever persistence needs — a conversation id, and the + ``not_in_memory`` flag some fixtures set — so a hand-built message becomes storable. + Storing is idempotent, so the helper is safe to apply to an already-persisted message. + + Args: + message (Message): The message to persist. + + Returns: + Message: The same message. + """ + memory = CentralMemory.get_memory_instance() + piece_ids = [piece.id for piece in message.message_pieces if piece.id is not None] + if not piece_ids or memory.get_message_pieces(prompt_ids=piece_ids): + return message + + conversation_id = next( + (piece.conversation_id for piece in message.message_pieces if piece.conversation_id), + str(uuid.uuid4()), + ) + for piece in message.message_pieces: + piece.not_in_memory = False + if not piece.conversation_id: + piece.conversation_id = conversation_id + + memory.add_message_to_memory(request=message) + return message + + def get_test_message_piece() -> MessagePiece: return MessagePiece( role="user", diff --git a/tests/unit/score/test_azure_content_filter.py b/tests/unit/score/test_azure_content_filter.py index 48a110c10d..072c8d42ba 100644 --- a/tests/unit/score/test_azure_content_filter.py +++ b/tests/unit/score/test_azure_content_filter.py @@ -8,7 +8,7 @@ import pytest from azure.ai.contentsafety.models import TextCategory -from unit.mocks import get_audio_message_piece, get_image_message_piece, get_test_message_piece +from unit.mocks import get_audio_message_piece, get_image_message_piece, get_test_message_piece, store_message from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface @@ -43,7 +43,7 @@ async def test_score_async_unsupported_data_type_returns_zero( # Unified FloatScaleScorer fallback: when all pieces are filtered out, return a single # Score(0.0) instead of an empty list (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 @@ -335,7 +335,7 @@ async def test_azure_content_filter_scorer_blocked_returns_one_score_per_categor ) message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) assert len(scores) == 2 assert {s.score_category[0] for s in scores} == {TextCategory.HATE.value, TextCategory.VIOLENCE.value} @@ -360,7 +360,7 @@ async def test_azure_content_filter_scorer_blocked_default_categories_returns_fo ) message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) assert len(scores) == 4 assert {s.score_category[0] for s in scores} == {c.value for c in TextCategory} diff --git a/tests/unit/score/test_conversation_history_scorer.py b/tests/unit/score/test_conversation_history_scorer.py index 4671196a6c..8a6b32647e 100644 --- a/tests/unit/score/test_conversation_history_scorer.py +++ b/tests/unit/score/test_conversation_history_scorer.py @@ -5,10 +5,12 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from unit.mocks import store_message from pyrit.memory import CentralMemory from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score from pyrit.score import ( + ContentScorable, MessageScorable, MessageScorer, Scorer, @@ -145,7 +147,7 @@ async def test_conversation_history_scorer_score_async_success(patch_central_dat mock_scorer.validate_return_scores = MagicMock() scorer = create_conversation_scorer(scorer=mock_scorer) - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(scores) == 1 result_score = scores[0] @@ -171,21 +173,15 @@ async def test_conversation_history_scorer_score_async_success(patch_central_dat async def test_conversation_history_scorer_conversation_not_found(patch_central_database): + """Loose content has no conversation behind it, so there is no history to score.""" mock_scorer = MagicMock(spec=SelfAskGeneralFloatScaleScorer) mock_scorer._validator = ScorerPromptValidator(supported_data_types=["text"]) scorer = create_conversation_scorer(scorer=mock_scorer) - nonexistent_conversation_id = str(uuid.uuid4()) - message_piece = MessagePiece( - role="assistant", - original_value="Test response", - conversation_id=nonexistent_conversation_id, - ) - message = MagicMock() - message.message_pieces = [message_piece] - - with pytest.raises(RuntimeError, match=f"Conversation with ID {nonexistent_conversation_id} not found in memory"): - await scorer.score_async(scorable=MessageScorable(message=message)) + # A MessageScorable cannot reach this guard: resolving it requires the pieces to be in + # memory, and then their conversation is there too. + with pytest.raises(RuntimeError, match="not found in memory"): + await scorer.score_async(scorable=ContentScorable(value="Test response")) async def test_conversation_history_scorer_filters_roles_correctly(patch_central_database): @@ -235,7 +231,7 @@ async def test_conversation_history_scorer_filters_roles_correctly(patch_central mock_scorer.validate_return_scores = MagicMock() scorer = create_conversation_scorer(scorer=mock_scorer) - await scorer.score_async(scorable=MessageScorable(message=message)) + await scorer.score_async(scorable=MessageScorable.from_message(message)) call_args = mock_scorer._score_async.call_args called_message = call_args.kwargs["message"] @@ -280,7 +276,7 @@ async def test_conversation_history_scorer_preserves_metadata(patch_central_data scorer = create_conversation_scorer(scorer=mock_scorer) - await scorer.score_async(scorable=MessageScorable(message=message)) + await scorer.score_async(scorable=MessageScorable.from_message(message)) call_args = mock_scorer._score_async.call_args called_message = call_args.kwargs["message"] @@ -331,7 +327,7 @@ async def test_conversation_scorer_persists_scores_exactly_once(patch_central_da conv_scorer = create_conversation_scorer(scorer=mock_scorer) message = MagicMock() message.message_pieces = [message_piece] - result_scores = await conv_scorer.score_async(scorable=MessageScorable(message=message)) + result_scores = await conv_scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(result_scores) == 1 assert result_scores[0].id == original_id, ( @@ -522,16 +518,9 @@ async def test_conversation_scorer_uses_partial_content_when_score_blocked_conte memory.add_message_pieces_to_memory(message_pieces=message_pieces) - # Use a text piece as the incoming message for validation purposes. - # ConversationScorer only uses it for conversation_id lookup — actual content comes from DB. - lookup_piece = MessagePiece( - role="assistant", - original_value="lookup", - conversation_id=conversation_id, - ) - message = MagicMock() - message.message_pieces = [lookup_piece] - message.get_piece.return_value = lookup_piece + # Name a piece that is already in the conversation. A scorable is a reference, so a + # synthetic lookup piece would have to be persisted and would then join the history. + message = blocked_piece.to_message() mock_scorer = MagicMock(spec=SelfAskGeneralFloatScaleScorer) mock_scorer._validator = ScorerPromptValidator(supported_data_types=["text"]) @@ -551,7 +540,7 @@ async def test_conversation_scorer_uses_partial_content_when_score_blocked_conte scorer = create_conversation_scorer(scorer=mock_scorer) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(scores) == 1 @@ -595,15 +584,9 @@ async def test_conversation_scorer_uses_error_json_when_score_blocked_content_di memory.add_message_pieces_to_memory(message_pieces=message_pieces) - # Use a text piece as the incoming message for validation purposes. - lookup_piece = MessagePiece( - role="assistant", - original_value="lookup", - conversation_id=conversation_id, - ) - message = MagicMock() - message.message_pieces = [lookup_piece] - message.get_piece.return_value = lookup_piece + # Name a piece that is already in the conversation. A scorable is a reference, so a + # synthetic lookup piece would have to be persisted and would then join the history. + message = blocked_piece.to_message() mock_scorer = MagicMock(spec=SelfAskGeneralFloatScaleScorer) mock_scorer._validator = ScorerPromptValidator(supported_data_types=["text"]) @@ -623,7 +606,7 @@ async def test_conversation_scorer_uses_error_json_when_score_blocked_content_di scorer = create_conversation_scorer(scorer=mock_scorer) # score_blocked_content defaults to False - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(scores) == 1 @@ -691,7 +674,7 @@ async def test_conversation_scorer_blocked_input_message_does_not_raise(patch_ce scorer = create_conversation_scorer(scorer=mock_scorer) # Must not raise — previously raised ValueError on the blocked piece. - scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(blocked_message))) assert len(scores) == 1 mock_scorer._score_async.assert_awaited_once() @@ -785,7 +768,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st inner_scorer = HarmfulContentDetector() scorer = create_conversation_scorer(scorer=inner_scorer) - scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(blocked_message))) assert len(scores) == 1 # Must be 1.0 (real score from prior turns), NOT 0.0 (fallback from rejected synthetic piece) diff --git a/tests/unit/score/test_float_scale_threshold_scorer.py b/tests/unit/score/test_float_scale_threshold_scorer.py index 69500aa662..f3449e0c1a 100644 --- a/tests/unit/score/test_float_scale_threshold_scorer.py +++ b/tests/unit/score/test_float_scale_threshold_scorer.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from unit.mocks import store_message from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score @@ -260,7 +261,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st ) blocked_message = Message(message_pieces=[blocked_piece]) - scores = await threshold_scorer.score_async(scorable=MessageScorable(message=blocked_message)) + scores = await threshold_scorer.score_async(scorable=MessageScorable.from_message(store_message(blocked_message))) assert len(scores) == 1 binary_score = scores[0] diff --git a/tests/unit/score/test_gandalf_scorer.py b/tests/unit/score/test_gandalf_scorer.py index 487ce087dd..78d678517d 100644 --- a/tests/unit/score/test_gandalf_scorer.py +++ b/tests/unit/score/test_gandalf_scorer.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import get_mock_target_identifier, store_message from pyrit.exceptions.exception_classes import PyritException from pyrit.memory.memory_interface import MemoryInterface @@ -64,7 +64,7 @@ async def test_gandalf_scorer_score( mocked_post.return_value = MagicMock(json=lambda: {"success": password_correct, "message": "Message"}) - scores = await scorer.score_async(scorable=MessageScorable(message=response)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) assert len(scores) == 1 assert scores[0].get_value() == password_correct @@ -99,7 +99,7 @@ async def test_gandalf_scorer_set_system_prompt( mocked_post.return_value = MagicMock(json=lambda: {"success": True, "message": "Message"}) - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) chat_target.set_system_prompt.assert_called_once() @@ -121,10 +121,11 @@ async def test_gandalf_scorer_adds_to_memory(mocked_post, level: GandalfLevel, s mocked_post.return_value = MagicMock(json=lambda: {"success": True, "message": "Message"}) - with patch.object(sqlite_instance, "get_message_pieces", return_value=[generated_request.message_pieces[0]]): + patched_pieces = [generated_request.message_pieces[0], response.message_pieces[0]] + with patch.object(sqlite_instance, "get_message_pieces", return_value=patched_pieces): scorer = GandalfScorer(level=level, chat_target=chat_target) - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) @pytest.mark.parametrize("level", [GandalfLevel.LEVEL_1, GandalfLevel.LEVEL_2, GandalfLevel.LEVEL_3]) @@ -140,7 +141,7 @@ async def test_gandalf_scorer_runtime_error_retries(level: GandalfLevel, sqlite_ scorer = GandalfScorer(level=level, chat_target=chat_target) with pytest.raises(PyritException, match="Error in scorer GandalfScorer"): - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) assert chat_target.send_prompt_async.call_count == 1 @@ -167,4 +168,4 @@ async def test_gandalf_scorer_wraps_httpx_error_as_pyrit_exception(mocked_post, scorer = GandalfScorer(level=GandalfLevel.LEVEL_1, chat_target=chat_target) with pytest.raises(PyritException, match="Error in scorer GandalfScorer"): - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) diff --git a/tests/unit/score/test_insecure_code_scorer.py b/tests/unit/score/test_insecure_code_scorer.py index 7120148841..78954b837b 100644 --- a/tests/unit/score/test_insecure_code_scorer.py +++ b/tests/unit/score/test_insecure_code_scorer.py @@ -4,6 +4,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from unit.mocks import store_message from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, SeedPrompt, UnvalidatedScore @@ -47,7 +48,7 @@ async def test_insecure_code_scorer_valid_response(mock_chat_target): message = MessagePiece(role="user", original_value="sample code").to_message() # Call the score_async method - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) # Assertions assert len(scores) == 1 @@ -70,7 +71,7 @@ async def test_insecure_code_scorer_invalid_json(mock_chat_target): message = MessagePiece(role="user", original_value="sample code").to_message() with pytest.raises(InvalidJsonException, match="Error in scorer InsecureCodeScorer.*Invalid JSON"): - await scorer.score_async(scorable=MessageScorable(message=message)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) # Ensure memory functions were not called mock_add_scores.assert_not_called() @@ -109,7 +110,7 @@ async def test_score_async_unsupported_data_type_returns_zero(mock_chat_target, # Unified FloatScaleScorer fallback: returns a single Score(0.0) when all pieces are filtered # out (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py index e12f9f9f47..7f263e7071 100644 --- a/tests/unit/score/test_message_scorer.py +++ b/tests/unit/score/test_message_scorer.py @@ -6,11 +6,10 @@ import pytest -from pyrit.memory import MemoryInterface +from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.score import ( ContentScorable, - MessageReferenceScorable, MessageScorable, MessageScorer, Scorable, @@ -70,58 +69,72 @@ def _build_score(self, message_piece: MessagePiece, objective: str | None) -> Sc def _assistant_message(value: str = "response", conversation_id: str | None = None) -> Message: - return MessagePiece( + """Return an assistant message that is already in memory, since a scorable names ids.""" + message = MessagePiece( role="assistant", original_value=value, conversation_id=conversation_id or str(uuid.uuid4()), ).to_message() + CentralMemory.get_memory_instance().add_message_to_memory(request=message) + return message + + +def _error_message() -> Message: + """Return a stored assistant message that carries a blocked error result.""" + message = MessagePiece( + role="assistant", + original_value="blocked", + original_value_data_type="error", + response_error="blocked", + conversation_id=str(uuid.uuid4()), + ).to_message() + CentralMemory.get_memory_instance().add_message_to_memory(request=message) + return message @pytest.mark.usefixtures("patch_central_database") class TestScorableResolution: """MessageScorer reduces every message-shaped scorable to a single Message.""" - async def test_message_scorable_is_scored_directly(self): + async def test_message_scorable_resolves_from_memory(self, sqlite_instance: MemoryInterface): + message = _assistant_message("stored response") scorer = RecordingScorer() - message = _assistant_message() - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(scores) == 1 - assert scorer.scored_messages == [message] + assert scorer.scored_messages[0].get_value() == "stored response" - async def test_message_reference_scorable_resolves_from_memory(self, sqlite_instance: MemoryInterface): + async def test_message_scorable_resolves_by_piece_id(self, sqlite_instance: MemoryInterface): message = _assistant_message("stored response") - sqlite_instance.add_message_to_memory(request=message) piece_id = message.get_piece().id scorer = RecordingScorer() - scores = await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(piece_id,))) + scores = await scorer.score_async(scorable=MessageScorable(message_piece_ids=(piece_id,))) assert len(scores) == 1 assert scorer.scored_messages[0].get_value() == "stored response" - async def test_message_reference_scorable_not_in_memory_raises(self): + async def test_message_scorable_not_in_memory_raises(self): scorer = RecordingScorer() missing_id = uuid.uuid4() with pytest.raises(ValueError, match="No message pieces found in memory"): - await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(missing_id,))) + await scorer.score_async(scorable=MessageScorable(message_piece_ids=(missing_id,))) - async def test_message_reference_scorable_partially_in_memory_raises(self, sqlite_instance: MemoryInterface): + async def test_message_scorable_partially_in_memory_raises(self, sqlite_instance: MemoryInterface): """A partial resolution is a caller error, so it must not be scored silently.""" stored = _assistant_message("stored response") - sqlite_instance.add_message_to_memory(request=stored) stored_id = stored.get_piece().id missing_id = uuid.uuid4() scorer = RecordingScorer() with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): - await scorer.score_async(scorable=MessageReferenceScorable(message_piece_ids=(stored_id, missing_id))) + await scorer.score_async(scorable=MessageScorable(message_piece_ids=(stored_id, missing_id))) assert scorer.scored_messages == [] - async def test_message_reference_scorable_spanning_messages_raises(self, sqlite_instance: MemoryInterface): + async def test_message_scorable_spanning_messages_raises(self, sqlite_instance: MemoryInterface): conversation_id = str(uuid.uuid4()) first = MessagePiece( role="user", original_value="ask", conversation_id=conversation_id, sequence=0 @@ -135,7 +148,7 @@ async def test_message_reference_scorable_spanning_messages_raises(self, sqlite_ with pytest.raises(ValueError, match="exactly one message"): await scorer.score_async( - scorable=MessageReferenceScorable( + scorable=MessageScorable( message_piece_ids=(first.get_piece().id, second.get_piece().id), ) ) @@ -185,7 +198,7 @@ async def test_role_filter_mismatch_skips_scoring(self): scorer = RecordingScorer() message = _assistant_message() - scores = await scorer.score_async(scorable=MessageScorable(message=message, role_filter="user")) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message, role_filter="user")) assert scores == [] assert scorer.scored_messages == [] @@ -194,34 +207,24 @@ async def test_role_filter_match_scores(self): scorer = RecordingScorer() message = _assistant_message() - scores = await scorer.score_async(scorable=MessageScorable(message=message, role_filter="assistant")) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message, role_filter="assistant")) assert len(scores) == 1 async def test_skip_on_error_result_skips_error_message(self): scorer = RecordingScorer() - message = MessagePiece( - role="assistant", - original_value="blocked", - original_value_data_type="error", - response_error="blocked", - ).to_message() + message = _error_message() - scores = await scorer.score_async(scorable=MessageScorable(message=message, skip_on_error_result=True)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message, skip_on_error_result=True)) assert scores == [] assert scorer.scored_messages == [] async def test_error_message_is_scored_when_not_skipping(self): scorer = RecordingScorer() - message = MessagePiece( - role="assistant", - original_value="blocked", - original_value_data_type="error", - response_error="blocked", - ).to_message() + message = _error_message() - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) assert len(scores) == 1 @@ -234,7 +237,7 @@ async def test_objective_reaches_the_scorer(self): scorer = RecordingScorer() await scorer.score_async( - scorable=MessageScorable(message=_assistant_message()), + scorable=MessageScorable.from_message(_assistant_message()), expectation=ScoringExpectation(objective="find the objective"), ) @@ -243,7 +246,7 @@ async def test_objective_reaches_the_scorer(self): async def test_no_expectation_means_no_objective(self): scorer = RecordingScorer() - await scorer.score_async(scorable=MessageScorable(message=_assistant_message())) + await scorer.score_async(scorable=MessageScorable.from_message(_assistant_message())) assert scorer.scored_objectives == [None] @@ -283,7 +286,6 @@ async def test_message_does_not_widen_to_the_stored_conversation(self, sqlite_in ).to_message() ) message = _assistant_message("only this turn", conversation_id=conversation_id) - sqlite_instance.add_message_to_memory(request=message) scorer = RecordingScorer() with pytest.warns(DeprecationWarning, match="Scorer.score_async"): @@ -309,12 +311,7 @@ async def test_legacy_role_filter_maps_onto_the_scorable(self): async def test_legacy_skip_on_error_result_maps_onto_the_scorable(self): scorer = RecordingScorer() - message = MessagePiece( - role="assistant", - original_value="blocked", - original_value_data_type="error", - response_error="blocked", - ).to_message() + message = _error_message() with pytest.warns(DeprecationWarning, match="Scorer.score_async"): scores = await scorer.score_async(message, skip_on_error_result=True) @@ -332,7 +329,6 @@ async def test_infer_objective_from_request_reads_the_previous_turn(self, sqlite ).to_message() ) message = _assistant_message("response", conversation_id=conversation_id) - sqlite_instance.add_message_to_memory(request=message) scorer = RecordingScorer() with pytest.warns(DeprecationWarning, match="Scorer.score_async"): @@ -344,7 +340,7 @@ async def test_new_signature_emits_no_warning(self, recwarn): scorer = RecordingScorer() await scorer.score_async( - scorable=MessageScorable(message=_assistant_message()), + scorable=MessageScorable.from_message(_assistant_message()), expectation=ScoringExpectation(objective="objective"), ) @@ -360,7 +356,7 @@ async def test_message_and_scorable_together_raises(self): message = _assistant_message() with pytest.raises(ValueError, match="not both"): - await scorer.score_async(message, scorable=MessageScorable(message=message)) + await scorer.score_async(message, scorable=MessageScorable.from_message(message)) async def test_neither_message_nor_scorable_raises(self): scorer = RecordingScorer() @@ -373,7 +369,7 @@ async def test_objective_and_expectation_together_raises(self): with pytest.raises(ValueError, match="not both"): await scorer.score_async( - scorable=MessageScorable(message=_assistant_message()), + scorable=MessageScorable.from_message(_assistant_message()), objective="one", expectation=ScoringExpectation(objective="two"), ) @@ -383,7 +379,7 @@ async def test_message_flags_with_a_scorable_raises(self, kwargs): scorer = RecordingScorer() with pytest.raises(ValueError, match="fields on the message scorable"): - await scorer.score_async(scorable=MessageScorable(message=_assistant_message()), **kwargs) + await scorer.score_async(scorable=MessageScorable.from_message(_assistant_message()), **kwargs) @pytest.mark.usefixtures("patch_central_database") @@ -398,7 +394,6 @@ def test_reads_the_turn_before_the_response(self, sqlite_instance: MemoryInterfa ).to_message() ) message = _assistant_message("the response", conversation_id=conversation_id) - sqlite_instance.add_message_to_memory(request=message) objective = extract_objective_from_previous_turn(message=message, memory=sqlite_instance) @@ -410,6 +405,8 @@ def test_returns_empty_for_a_user_message(self, sqlite_instance: MemoryInterface assert extract_objective_from_previous_turn(message=message, memory=sqlite_instance) == "" def test_returns_empty_when_the_conversation_is_not_stored(self, sqlite_instance: MemoryInterface): - message = _assistant_message() + message = MessagePiece( + role="assistant", original_value="a response", conversation_id=str(uuid.uuid4()) + ).to_message() assert extract_objective_from_previous_turn(message=message, memory=sqlite_instance) == "" diff --git a/tests/unit/score/test_plagiarism_scorer.py b/tests/unit/score/test_plagiarism_scorer.py index eed3048b50..382369c10b 100644 --- a/tests/unit/score/test_plagiarism_scorer.py +++ b/tests/unit/score/test_plagiarism_scorer.py @@ -1,12 +1,12 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest +from unit.mocks import mock_memory_resolving, store_message from pyrit.memory import CentralMemory -from pyrit.memory.memory_interface import MemoryInterface from pyrit.models import MessagePiece from pyrit.score import MessageScorable, PlagiarismMetric, PlagiarismScorer @@ -52,7 +52,7 @@ async def test_score_async_lcs_metric(self): request = message_piece.to_message() - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 score = scores[0] @@ -170,7 +170,6 @@ async def test_score_async_completely_different_texts(self): async def test_score_async_adds_to_memory(self): """Test that scoring adds results to memory.""" - memory = MagicMock(MemoryInterface) reference_text = "Test reference text" scorer = PlagiarismScorer(reference_text=reference_text) @@ -181,8 +180,9 @@ async def test_score_async_adds_to_memory(self): converted_value_data_type="text", ).to_message() + memory = mock_memory_resolving(request) with patch.object(CentralMemory, "get_memory_instance", return_value=memory): - await scorer.score_async(scorable=MessageScorable(message=request)) + await scorer.score_async(scorable=MessageScorable.from_message(request)) memory.add_scores_to_memory.assert_called_once() async def test_score_async_unsupported_data_type_returns_zero(self, patch_central_database): @@ -199,7 +199,7 @@ async def test_score_async_unsupported_data_type_returns_zero(self, patch_centra # Unified FloatScaleScorer fallback: returns a single Score(0.0) when all pieces are filtered # out (mirrors TrueFalseScorer's no-pieces fallback). - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].score_type == "float_scale" assert scores[0].get_value() == 0.0 diff --git a/tests/unit/score/test_question_answer_scorer.py b/tests/unit/score/test_question_answer_scorer.py index 2654c6b452..c41c0ae04d 100644 --- a/tests/unit/score/test_question_answer_scorer.py +++ b/tests/unit/score/test_question_answer_scorer.py @@ -3,13 +3,12 @@ import os import uuid -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest -from unit.mocks import get_image_message_piece +from unit.mocks import get_image_message_piece, mock_memory_resolving, store_message from pyrit.memory.central_memory import CentralMemory -from pyrit.memory.memory_interface import MemoryInterface from pyrit.models import Message, MessagePiece from pyrit.score import MessageScorable, QuestionAnswerScorer @@ -39,7 +38,7 @@ async def test_score_async_unsupported_image_type_returns_false( message = Message(message_pieces=[image_message_piece]) # With raise_on_no_valid_pieces=False (default), returns False for unsupported data types - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale @@ -58,7 +57,7 @@ async def test_score_async_missing_metadata_returns_false(patch_central_database scorer = QuestionAnswerScorer(category=["new_category"]) # With raise_on_no_valid_pieces=False (default), returns False for missing metadata - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale @@ -80,7 +79,7 @@ async def test_question_answer_scorer_score(response: str, expected_score: bool, scorer = QuestionAnswerScorer(category=["new_category"]) message = Message(message_pieces=[text_message_piece]) - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) assert len(scores) == 1 result_score = scores[0] @@ -90,33 +89,33 @@ async def test_question_answer_scorer_score(response: str, expected_score: bool, async def test_question_answer_scorer_adds_to_memory(): - memory = MagicMock(MemoryInterface) + message = MessagePiece( + role="user", + original_value="test content", + converted_value="0: Paris", + converted_value_data_type="text", + prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, + ).to_message() + memory = mock_memory_resolving(message) with patch.object(CentralMemory, "get_memory_instance", return_value=memory): scorer = QuestionAnswerScorer(category=["new_category"]) - message = MessagePiece( - role="user", - original_value="test content", - converted_value="0: Paris", - converted_value_data_type="text", - prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, - ).to_message() - await scorer.score_async(scorable=MessageScorable(message=message)) + await scorer.score_async(scorable=MessageScorable.from_message(message)) memory.add_scores_to_memory.assert_called_once() async def test_question_answer_scorer_no_category(): - memory = MagicMock(MemoryInterface) + message = MessagePiece( + role="user", + original_value="test content", + converted_value="0: Paris", + converted_value_data_type="text", + prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, + ).to_message() + memory = mock_memory_resolving(message) with patch.object(CentralMemory, "get_memory_instance", return_value=memory): scorer = QuestionAnswerScorer() - message = MessagePiece( - role="user", - original_value="test content", - converted_value="0: Paris", - converted_value_data_type="text", - prompt_metadata={"correct_answer_index": "0", "correct_answer": "Paris"}, - ).to_message() - await scorer.score_async(scorable=MessageScorable(message=message)) + await scorer.score_async(scorable=MessageScorable.from_message(message)) memory.add_scores_to_memory.assert_called_once() diff --git a/tests/unit/score/test_scorable.py b/tests/unit/score/test_scorable.py index c3ab4c8c71..369d5774ba 100644 --- a/tests/unit/score/test_scorable.py +++ b/tests/unit/score/test_scorable.py @@ -3,21 +3,16 @@ import dataclasses import uuid -from unittest.mock import MagicMock import pytest from pyrit.memory import MemoryInterface -from pyrit.models import MessagePiece -from pyrit.score import ContentScorable, MessageReferenceScorable, MessageScorable, Scorable, SingleMessageScorable - - -def _message(): - return MessagePiece(role="assistant", original_value="hello").to_message() +from pyrit.models import Message, MessagePiece +from pyrit.score import ContentScorable, MessageScorable, Scorable def _stored_message(value: str = "stored response"): - """Return a message that memory will accept, so it needs a conversation id.""" + """Return a message memory will accept, so it needs a conversation id.""" return MessagePiece( role="assistant", original_value=value, @@ -25,16 +20,10 @@ def _stored_message(value: str = "stored response"): ).to_message() -def _no_memory() -> MemoryInterface: - """Return a memory that the test can assert was never touched.""" - return MagicMock(spec=MemoryInterface) - - @pytest.mark.parametrize( "scorable, field_name", [ - (MessageScorable(message=_message()), "message"), - (MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), + (MessageScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), (ContentScorable(value="hello"), "value"), ], ) @@ -46,21 +35,21 @@ def test_scorable_is_frozen(scorable, field_name): @pytest.mark.parametrize( "scorable", [ - MessageScorable(message=_message()), - MessageReferenceScorable(message_piece_ids=(uuid.uuid4(),)), + MessageScorable(message_piece_ids=(uuid.uuid4(),)), + ContentScorable(value="hello"), ], ) -def test_scorables_naming_a_message_share_one_base(scorable): - assert isinstance(scorable, SingleMessageScorable) +def test_every_scorable_is_a_scorable(scorable): assert isinstance(scorable, Scorable) def test_content_scorable_names_content_rather_than_a_message(): - """Loose content has no message behind it, so it must not inherit message identity.""" + """Loose content has no message behind it, so it must not answer the message contract.""" scorable = ContentScorable(value="hello") assert isinstance(scorable, Scorable) - assert not isinstance(scorable, SingleMessageScorable) + assert not isinstance(scorable, MessageScorable) + assert not hasattr(scorable, "resolve_message") @pytest.mark.parametrize("field_name", ["role_filter", "skip_on_error_result"]) @@ -71,55 +60,40 @@ def test_content_scorable_rejects_message_filters(field_name): ContentScorable(value="hello", **{field_name: "assistant"}) -def test_single_message_scorable_cannot_be_instantiated(): - # The base only declares the contract; a scorable must say how it resolves. - assert "resolve_message" in SingleMessageScorable.__abstractmethods__ - - with pytest.raises(TypeError): - SingleMessageScorable() # type: ignore[abstract] - - def test_scorables_are_keyword_only(): with pytest.raises(TypeError): ContentScorable("hello") # type: ignore[misc] class TestMessageScorable: - def test_carries_the_message_and_filters(self): - message = _message() - scorable = MessageScorable(message=message, role_filter="assistant", skip_on_error_result=True) - - assert scorable.message is message - assert scorable.role_filter == "assistant" - assert scorable.skip_on_error_result is True - - def test_resolves_to_the_message_it_holds_without_memory(self): - message = _message() - memory = _no_memory() - - assert MessageScorable(message=message).resolve_message(memory=memory) is message - memory.get_message_pieces.assert_not_called() # type: ignore[attr-defined] - - -class TestMessageReferenceScorable: def test_defaults(self): piece_id = uuid.uuid4() - scorable = MessageReferenceScorable(message_piece_ids=(piece_id,)) + scorable = MessageScorable(message_piece_ids=(piece_id,)) assert scorable.message_piece_ids == (piece_id,) assert scorable.role_filter is None assert scorable.skip_on_error_result is False + def test_from_message_names_the_pieces_rather_than_carrying_them(self): + """A scorable is a reference, so a message in hand becomes its piece ids.""" + message = _stored_message() + + scorable = MessageScorable.from_message(message, role_filter="assistant") + + assert scorable.message_piece_ids == (message.get_piece().id,) + assert scorable.role_filter == "assistant" + assert not hasattr(scorable, "message") + def test_resolves_from_memory(self, sqlite_instance: MemoryInterface): stored = _stored_message() sqlite_instance.add_message_to_memory(request=stored) - scorable = MessageReferenceScorable(message_piece_ids=(stored.get_piece().id,)) + scorable = MessageScorable.from_message(stored) assert scorable.resolve_message(memory=sqlite_instance).get_value() == "stored response" def test_raises_when_nothing_is_in_memory(self, sqlite_instance: MemoryInterface): missing_id = uuid.uuid4() - scorable = MessageReferenceScorable(message_piece_ids=(missing_id,)) + scorable = MessageScorable(message_piece_ids=(missing_id,)) with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): scorable.resolve_message(memory=sqlite_instance) @@ -128,9 +102,8 @@ def test_raises_naming_only_the_missing_ids(self, sqlite_instance: MemoryInterfa """A partial resolution is a caller error, and the error must point at the bad ids.""" stored = _stored_message() sqlite_instance.add_message_to_memory(request=stored) - stored_id = stored.get_piece().id missing_id = uuid.uuid4() - scorable = MessageReferenceScorable(message_piece_ids=(stored_id, missing_id)) + scorable = MessageScorable(message_piece_ids=(stored.get_piece().id, missing_id)) with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): scorable.resolve_message(memory=sqlite_instance) @@ -145,9 +118,7 @@ def test_raises_when_pieces_span_more_than_one_message(self, sqlite_instance: Me ).to_message() sqlite_instance.add_message_to_memory(request=first) sqlite_instance.add_message_to_memory(request=second) - scorable = MessageReferenceScorable( - message_piece_ids=(first.get_piece().id, second.get_piece().id), - ) + scorable = MessageScorable(message_piece_ids=(first.get_piece().id, second.get_piece().id)) with pytest.raises(ValueError, match="exactly one message"): scorable.resolve_message(memory=sqlite_instance) @@ -171,6 +142,34 @@ def test_adapts_non_text_data_types(self): assert message.get_piece().original_value_data_type == "image_path" - def test_does_not_resolve_a_message(self): - # It has no message to name, so it must not answer the resolution contract. - assert not hasattr(ContentScorable(value="hello"), "resolve_message") + +class TestFromMessage: + """A caller holding a persisted Message names its pieces rather than carrying it.""" + + def test_persisted_message_becomes_a_reference(self): + message = _stored_message() + + scorable = MessageScorable.from_message(message) + + assert isinstance(scorable, MessageScorable) + assert scorable.message_piece_ids == (message.get_piece().id,) + + def test_filters_travel_onto_the_reference(self): + scorable = MessageScorable.from_message(_stored_message(), role_filter="assistant", skip_on_error_result=True) + + assert isinstance(scorable, MessageScorable) + assert scorable.role_filter == "assistant" + assert scorable.skip_on_error_result is True + + +def test_resolution_preserves_the_order_the_scorable_names(sqlite_instance: MemoryInterface): + """Memory returns its own order, but a multi-piece message reads differently if shuffled.""" + conversation_id = str(uuid.uuid4()) + first = MessagePiece(role="assistant", original_value="one", conversation_id=conversation_id, sequence=0) + second = MessagePiece(role="assistant", original_value="two", conversation_id=conversation_id, sequence=0) + sqlite_instance.add_message_to_memory(request=Message(message_pieces=[first, second])) + + reversed_scorable = MessageScorable(message_piece_ids=(second.id, first.id)) + + resolved = reversed_scorable.resolve_message(memory=sqlite_instance) + assert [piece.original_value for piece in resolved.message_pieces] == ["two", "one"] diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index d453fdd54b..2a85d84fe2 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -7,10 +7,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import get_mock_target_identifier, store_message from pyrit.exceptions import InvalidJsonException, remove_markdown_json -from pyrit.memory import MemoryInterface +from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.prompt_target import PromptTarget from pyrit.score import ( @@ -456,7 +456,7 @@ async def test_scorer_score_responses_batch_async(patch_central_database): # Get the call_args for the first call _, first_call_kwargs = mock_score_async.call_args_list[0] - assert first_call_kwargs["scorable"] == MessageScorable(message=user_req) + assert first_call_kwargs["scorable"] == MessageScorable.from_message(store_message(user_req)) assert first_call_kwargs["expectation"] == ScoringExpectation(objective="") assert fake_scores[0] in results @@ -508,17 +508,17 @@ async def test_score_image_batch_async_works_when_objectives_none(patch_central_ assert "objective" not in call_kwargs -async def test_score_response_async_empty_scorers(): +async def test_score_response_async_empty_scorers(patch_central_database): """Test that score_response_async returns empty list when no scorers provided.""" response = Message( message_pieces=[MessagePiece(role="assistant", original_value="test", conversation_id="test-convo")] ) - result = await Scorer.score_response_async(response=response, objective="test task") + result = await Scorer.score_response_async(response=store_message(response), objective="test task") assert result == {"auxiliary_scores": [], "objective_scores": []} -async def test_score_response_async_no_matching_role(): +async def test_score_response_async_no_matching_role(patch_central_database): """Test that score_response_async returns empty list when no pieces match role filter.""" response = Message( message_pieces=[ @@ -531,7 +531,7 @@ async def test_score_response_async_no_matching_role(): scorer.score_async = AsyncMock(return_value=[]) result = await Scorer.score_response_async( - response=response, + response=store_message(response), objective_scorer=scorer, auxiliary_scorers=[scorer], role_filter="assistant", @@ -541,7 +541,7 @@ async def test_score_response_async_no_matching_role(): scorer.score_async.assert_called() -async def test_score_response_async_parallel_execution(): +async def test_score_response_async_parallel_execution(patch_central_database): """Test that score_response_async runs all scorers in parallel on all filtered pieces.""" piece1 = MessagePiece(role="assistant", original_value="response1", conversation_id="test-convo") piece2 = MessagePiece(role="assistant", original_value="response2", conversation_id="test-convo") @@ -568,7 +568,9 @@ async def test_score_response_async_parallel_execution(): assert score1_1 in result["auxiliary_scores"] assert score2_1 in result["auxiliary_scores"] - expected_scorable = MessageScorable(message=response, role_filter="assistant", skip_on_error_result=True) + expected_scorable = MessageScorable.from_message( + store_message(response), role_filter="assistant", skip_on_error_result=True + ) scorer1.score_async.assert_any_call( scorable=expected_scorable, expectation=ScoringExpectation(objective="test task"), @@ -579,23 +581,25 @@ async def test_score_response_async_parallel_execution(): ) -async def test_score_response_select_first_success_async_empty_scorers(): +async def test_score_response_select_first_success_async_empty_scorers(patch_central_database): """Test that score_response_select_first_success_async returns None when no scorers provided.""" response = Message( message_pieces=[MessagePiece(role="assistant", original_value="test", conversation_id="test-convo")] ) - result = await Scorer.score_response_multiple_scorers_async(response=response, scorers=[], objective="test task") + result = await Scorer.score_response_multiple_scorers_async( + response=store_message(response), scorers=[], objective="test task" + ) assert result == [] -async def test_score_async_no_matching_role(): +async def test_score_async_no_matching_role(patch_central_database): """Test that score_response_select_first_success_async returns None when no pieces match role filter.""" response = Message(message_pieces=[MessagePiece(role="user", original_value="test", conversation_id="test-convo")]) scorer = MockScorer() result = await scorer.score_async( - scorable=MessageScorable(message=response, role_filter="assistant"), + scorable=MessageScorable.from_message(store_message(response), role_filter="assistant"), expectation=ScoringExpectation(objective="test task"), ) @@ -678,24 +682,28 @@ async def test_score_response_success_async_no_success_returns_first(): assert scorer2.score_async.call_count == 1 -async def test_score_response_success_async_parallel_scoring_per_piece(): +async def test_score_response_success_async_parallel_scoring_per_piece(patch_central_database): """Test that score_response_success_async runs scorers in parallel for each piece.""" piece1 = MessagePiece(role="assistant", original_value="response1", conversation_id="test-convo") piece2 = MessagePiece(role="assistant", original_value="response2", conversation_id="test-convo") - response = Message(message_pieces=[piece1, piece2]) + response = store_message(Message(message_pieces=[piece1, piece2])) # Track call order call_order = [] + def _first_value(scorable: MessageScorable) -> str: + # A scorable names pieces rather than carrying them, so read it back from memory. + return scorable.resolve_message(memory=CentralMemory.get_memory_instance()).message_pieces[0].original_value + async def mock_score_async_1(*, scorable: MessageScorable, **kwargs) -> list[Score]: - call_order.append(("scorer1", scorable.message.message_pieces[0].original_value)) + call_order.append(("scorer1", _first_value(scorable))) score = MagicMock(spec=Score) score.get_value.return_value = False return [score] async def mock_score_async_2(*, scorable: MessageScorable, **kwargs) -> list[Score]: - call_order.append(("scorer2", scorable.message.message_pieces[0].original_value)) + call_order.append(("scorer2", _first_value(scorable))) score = MagicMock(spec=Score) score.get_value.return_value = False return [score] @@ -805,7 +813,7 @@ async def test_score_response_async_both_types(): assert result["objective_scores"][0] == obj_score -async def test_score_response_async_multiple_pieces(): +async def test_score_response_async_multiple_pieces(patch_central_database): """Test score_response_async with multiple response pieces.""" piece1 = MessagePiece(role="assistant", original_value="response1", conversation_id="test-convo") piece2 = MessagePiece(role="assistant", original_value="response2", conversation_id="test-convo") @@ -828,7 +836,7 @@ async def test_score_response_async_multiple_pieces(): obj_scorer.score_async = AsyncMock(return_value=[obj_score]) result = await Scorer.score_response_async( - response=response, + response=store_message(response), auxiliary_scorers=[aux_scorer1, aux_scorer2], objective_scorer=obj_scorer, objective="test task", @@ -848,7 +856,7 @@ async def test_score_response_async_multiple_pieces(): assert result["objective_scores"][0] == obj_score -async def test_score_response_async_skip_on_error_true(): +async def test_score_response_async_skip_on_error_true(patch_central_database): """Test score_response_async skips error pieces when skip_on_error_result=True.""" piece1 = MessagePiece(role="assistant", original_value="good response", conversation_id="test-convo") piece2 = MessagePiece( @@ -869,7 +877,7 @@ async def test_score_response_async_skip_on_error_true(): obj_scorer.score_async = AsyncMock(return_value=[obj_score]) result = await Scorer.score_response_async( - response=response, + response=store_message(response), auxiliary_scorers=[aux_scorer], objective_scorer=obj_scorer, objective="test task", @@ -885,7 +893,7 @@ async def test_score_response_async_skip_on_error_true(): obj_scorer.score_async.assert_called_once() -async def test_score_response_async_skip_on_error_false(): +async def test_score_response_async_skip_on_error_false(patch_central_database): """Test score_response_async includes error pieces when skip_on_error_result=False.""" piece1 = MessagePiece(role="assistant", original_value="good response", conversation_id="test-convo") piece2 = MessagePiece( @@ -906,7 +914,7 @@ async def test_score_response_async_skip_on_error_false(): obj_scorer.score_async = AsyncMock(return_value=[obj_score]) result = await Scorer.score_response_async( - response=response, + response=store_message(response), auxiliary_scorers=[aux_scorer], objective_scorer=obj_scorer, objective="test task", @@ -1050,7 +1058,7 @@ async def test_get_supported_pieces_filters_unsupported_data_types(patch_central response = Message(message_pieces=[text_piece, image_piece, audio_piece]) # Score the response - scores = await scorer.score_async(scorable=MessageScorable(message=response)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) # Should only score the text piece assert len(scorer.scored_piece_ids) == 1 @@ -1084,7 +1092,7 @@ async def test_unsupported_pieces_ignored_when_enforce_all_pieces_valid_false(pa response = Message(message_pieces=[image_piece, text_piece]) # Should not raise an error, just skip the image piece - scores = await scorer.score_async(scorable=MessageScorable(message=response)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) assert len(scores) == 1 assert len(scorer.scored_piece_ids) == 1 @@ -1116,7 +1124,7 @@ async def test_all_unsupported_pieces_raises_error(patch_central_database): # Should raise error from validator because no valid pieces to score with pytest.raises(ValueError, match="There are no valid pieces to score"): - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) # No pieces should have been scored assert len(scorer.scored_piece_ids) == 0 @@ -1173,7 +1181,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st response = Message(message_pieces=[text_piece, image_piece]) # Score the response - scores = await scorer.score_async(scorable=MessageScorable(message=response)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) # Should only score the text piece assert len(scorer.scored_piece_ids) == 1 @@ -1209,7 +1217,7 @@ async def test_base_scorer_score_async_implementation(patch_central_database): response = Message(message_pieces=[text_piece1, text_piece2]) # Score the response - scores = await scorer.score_async(scorable=MessageScorable(message=response)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) # Should score both pieces assert len(scorer.scored_piece_ids) == 2 @@ -1341,7 +1349,9 @@ async def test_blocked_response_returns_specific_rationale( ) response = Message(message_pieces=[blocked_piece]) - scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await true_false_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1363,7 +1373,9 @@ async def test_error_response_returns_specific_rationale( ) response = Message(message_pieces=[error_piece]) - scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await true_false_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1385,7 +1397,9 @@ async def test_filtered_pieces_returns_generic_rationale( ) response = Message(message_pieces=[normal_piece]) - scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await true_false_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1408,7 +1422,9 @@ async def test_blocked_takes_precedence_over_generic_error( ) response = Message(message_pieces=[blocked_piece]) - scores = await true_false_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await true_false_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) # Should specifically mention blocked, not generic error assert "blocked" in scores[0].score_rationale.lower() @@ -1467,7 +1483,9 @@ async def test_blocked_response_returns_zero_with_blocked_rationale( ) response = Message(message_pieces=[blocked_piece]) - scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await float_scale_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].score_type == "float_scale" @@ -1489,7 +1507,9 @@ async def test_other_error_response_returns_zero_with_error_rationale( ) response = Message(message_pieces=[error_piece]) - scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await float_scale_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].get_value() == 0.0 @@ -1510,7 +1530,9 @@ async def test_filtered_pieces_return_zero_with_generic_rationale( ) response = Message(message_pieces=[normal_piece]) - scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await float_scale_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) assert len(scores) == 1 assert scores[0].get_value() == 0.0 @@ -1536,7 +1558,9 @@ async def test_text_only_scorer_filters_blocked_via_validator( with patch.object( float_scale_scorer_returns_empty, "_score_piece_async", new_callable=AsyncMock ) as mock_score_piece: - scores = await float_scale_scorer_returns_empty.score_async(scorable=MessageScorable(message=response)) + scores = await float_scale_scorer_returns_empty.score_async( + scorable=MessageScorable.from_message(store_message(response)) + ) mock_score_piece.assert_not_called() assert len(scores) == 1 @@ -1818,14 +1842,16 @@ async def test_raises_by_default(self): scorer = _ForwarderTrueFalseScorer(chat_target=_make_scorer_blocking_target()) with pytest.raises(ScorerLLMResponseBlockedException, match="blocked by content filtering"): - await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(_make_normal_input_message()))) async def test_returns_false_when_flag_disabled(self): target = _make_scorer_blocking_target() scorer = _ForwarderTrueFalseScorer(chat_target=target) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(_make_normal_input_message())) + ) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -1837,7 +1863,9 @@ async def test_returns_zero_for_float_scale_when_flag_disabled(self): scorer = _ForwarderFloatScaleScorer(chat_target=_make_scorer_blocking_target()) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(_make_normal_input_message())) + ) assert len(scores) == 1 assert scores[0].score_value == "0.0" @@ -1849,13 +1877,15 @@ async def test_direct_transport_caller_raises_by_default(self): scorer = _DirectTransportTrueFalseScorer(chat_target=_make_scorer_blocking_target()) with pytest.raises(ScorerLLMResponseBlockedException, match="blocked by content filtering"): - await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(_make_normal_input_message()))) async def test_direct_transport_caller_returns_false_when_flag_disabled(self): scorer = _DirectTransportTrueFalseScorer(chat_target=_make_scorer_blocking_target()) scorer.raise_if_scorer_blocks = False - scores = await scorer.score_async(scorable=MessageScorable(message=_make_normal_input_message())) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(_make_normal_input_message())) + ) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2059,7 +2089,7 @@ async def test_default_false_skips_blocked_piece_text_only_scorer(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2071,7 +2101,7 @@ async def test_true_substitutes_blocked_piece_for_text_only_scorer(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2084,7 +2114,7 @@ async def test_refusal_scorer_short_circuits_on_blocked_by_default(self): scorer = _MockRefusalScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2096,7 +2126,7 @@ async def test_refusal_scorer_evaluates_partial_content_when_flag_on(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2109,7 +2139,7 @@ async def test_no_substitute_when_no_partial_content(self): msg = Message(message_pieces=[_make_blocked_piece()]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 assert scores[0].score_value == "false" @@ -2120,10 +2150,10 @@ async def test_normal_piece_unaffected_by_flag(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_normal_piece()]) - scores_off = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores_off = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) scorer.scored_pieces.clear() scorer.score_blocked_content = True - scores_on = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores_on = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert scores_off[0].score_value == scores_on[0].score_value @@ -2133,7 +2163,7 @@ async def test_mixed_pieces_only_blocked_substituted(self): msg = Message(message_pieces=[_make_normal_piece(), _make_blocked_piece(partial_content="partial harmful")]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(msg))) assert len(scores) == 1 # TrueFalseScorer aggregates assert len(scorer.scored_pieces) == 2 @@ -2151,7 +2181,9 @@ async def test_skip_on_error_true_without_flag_skips_blocked(self): scorer = _BlockedContentScorer() msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert scores == [] async def test_skip_on_error_true_with_flag_does_not_skip_when_partial_content(self): @@ -2159,7 +2191,9 @@ async def test_skip_on_error_true_with_flag_does_not_skip_when_partial_content(s msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2168,7 +2202,9 @@ async def test_skip_on_error_true_with_flag_still_skips_when_no_partial_content( msg = Message(message_pieces=[_make_blocked_piece()]) scorer.score_blocked_content = True - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert scores == [] async def test_skip_on_error_skips_error_type_without_response_error_flag(self): @@ -2185,7 +2221,9 @@ async def test_skip_on_error_skips_error_type_without_response_error_flag(self): ] ) - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert scores == [] assert scorer.scored_pieces == [] @@ -2203,7 +2241,9 @@ async def test_skip_on_error_scores_structured_refusal_as_text(self, validator: piece = _make_blocked_piece(structured_refusal=refusal) msg = Message(message_pieces=[piece]) - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert len(scores) == 1 assert scorer.scored_pieces[0].id == piece.id @@ -2231,7 +2271,9 @@ async def test_skip_on_error_still_skips_mixed_structured_and_runtime_errors(sel ] ) - scores = await scorer.score_async(scorable=MessageScorable(message=msg, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + ) assert scores == [] assert scorer.scored_pieces == [] @@ -2248,7 +2290,7 @@ async def test_score_response_async_passes_flag_to_scorers(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) result = await Scorer.score_response_async( - response=msg, + response=store_message(msg), objective_scorer=obj_scorer, objective="test", skip_on_error_result=False, @@ -2263,7 +2305,7 @@ async def test_score_response_async_default_does_not_substitute(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) result = await Scorer.score_response_async( - response=msg, + response=store_message(msg), objective_scorer=obj_scorer, objective="test", skip_on_error_result=False, @@ -2280,7 +2322,7 @@ async def test_score_response_multiple_scorers_passes_flag(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scores = await Scorer.score_response_multiple_scorers_async( - response=msg, + response=store_message(msg), scorers=[scorer1, scorer2], objective="test", skip_on_error_result=False, diff --git a/tests/unit/score/test_self_ask_category.py b/tests/unit/score/test_self_ask_category.py index cbbc9e962c..69fd14e76a 100644 --- a/tests/unit/score/test_self_ask_category.py +++ b/tests/unit/score/test_self_ask_category.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import get_mock_target_identifier, store_message from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory @@ -256,15 +256,25 @@ async def test_score_prompts_batch_async( chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget") chat_target.send_prompt_async = AsyncMock() chat_target._max_requests_per_minute = max_requests_per_minute - with patch.object(CentralMemory, "get_memory_instance", return_value=MagicMock()): + + prompt = MessagePiece(role="assistant", original_value="test").to_message() + prompt2 = MessagePiece(role="assistant", original_value="test 2").to_message() + + # Scoring resolves a scorable through memory, so the fake has to answer id lookups. + # A real database is not wanted here: the scorer would persist the same mocked + # response twice and collide on its primary key. + known = {str(piece.id): piece for message in (prompt, prompt2) for piece in message.message_pieces} + memory = MagicMock() + memory.get_message_pieces.side_effect = lambda **kwargs: [ + known[str(piece_id)] for piece_id in kwargs.get("prompt_ids", []) if str(piece_id) in known + ] + + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): scorer = SelfAskCategoryScorer.from_content_classifier( chat_target=chat_target, content_classifier=HARM_CLASSIFIER, ) - prompt = MessagePiece(role="assistant", original_value="test").to_message() - prompt2 = MessagePiece(role="assistant", original_value="test 2").to_message() - with patch.object(chat_target, "send_prompt_async", return_value=[scorer_category_response_false]): if batch_size != 1 and max_requests_per_minute: with pytest.raises(ValueError): @@ -299,7 +309,7 @@ async def test_blocked_response_returns_false_without_invoking_llm(patch_central ) blocked_message = Message(message_pieces=[blocked_piece]) - scores = await scorer.score_async(scorable=MessageScorable(message=blocked_message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(blocked_message))) chat_target.send_prompt_async.assert_not_called() assert len(scores) == 1 diff --git a/tests/unit/score/test_self_ask_question_answer_scorer.py b/tests/unit/score/test_self_ask_question_answer_scorer.py index 09024ef381..6c846ec418 100644 --- a/tests/unit/score/test_self_ask_question_answer_scorer.py +++ b/tests/unit/score/test_self_ask_question_answer_scorer.py @@ -4,6 +4,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from unit.mocks import store_message from pyrit.models import ComponentIdentifier, MessagePiece, Score, ScoringExpectation, UnvalidatedScore from pyrit.prompt_target import PromptTarget @@ -42,7 +43,8 @@ async def test_score_async_returns_score_from_unvalidated(mock_chat_target): new=AsyncMock(return_value=unvalidated), ): scores = await scorer.score_async( - scorable=MessageScorable(message=message), expectation=ScoringExpectation(objective="2+2=?\nanswer: 4") + scorable=MessageScorable.from_message(store_message(message)), + expectation=ScoringExpectation(objective="2+2=?\nanswer: 4"), ) assert len(scores) == 1 diff --git a/tests/unit/score/test_self_ask_refusal.py b/tests/unit/score/test_self_ask_refusal.py index 50a43bd2bf..78d8054537 100644 --- a/tests/unit/score/test_self_ask_refusal.py +++ b/tests/unit/score/test_self_ask_refusal.py @@ -8,7 +8,7 @@ from uuid import uuid4 import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import get_mock_target_identifier, store_message from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory @@ -254,7 +254,7 @@ async def test_score_async_filtered_response(patch_central_database): conversation_id=str(uuid4()), ).to_message() memory.add_message_pieces_to_memory(message_pieces=request.message_pieces) - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].score_value == "true" diff --git a/tests/unit/score/test_shieldgemma_scorer.py b/tests/unit/score/test_shieldgemma_scorer.py index 6c1aef2438..cf27cebb90 100644 --- a/tests/unit/score/test_shieldgemma_scorer.py +++ b/tests/unit/score/test_shieldgemma_scorer.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import get_mock_target_identifier, store_message from pyrit.exceptions import InvalidJsonException from pyrit.memory.memory_interface import MemoryInterface @@ -143,7 +143,7 @@ async def test_response_scoring_excludes_a_stored_user_turn(sqlite_instance: Mem target = _mock_target("No") scorer = ShieldGemmaScorer(chat_target=target, guideline=CUSTOM_GUIDELINE) - await scorer.score_async(scorable=MessageScorable(message=response)) + await scorer.score_async(scorable=MessageScorable.from_message(store_message(response))) sent = _sent_request(target) assert "Chatbot Response: A response judged on its own." in sent @@ -232,7 +232,7 @@ async def test_multiple_pieces_keep_every_verdict_and_report_the_aggregate( ) message.set_response_not_in_memory() - scores = await scorer.score_async(scorable=MessageScorable(message=message)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(message))) assert target.send_prompt_async.call_count == 2 assert scores[0].get_value() is True diff --git a/tests/unit/score/test_substring.py b/tests/unit/score/test_substring.py index 6f65f024aa..1bf311c371 100644 --- a/tests/unit/score/test_substring.py +++ b/tests/unit/score/test_substring.py @@ -5,7 +5,7 @@ from unittest.mock import MagicMock, patch import pytest -from unit.mocks import get_image_message_piece +from unit.mocks import get_image_message_piece, store_message from pyrit.analytics import ApproximateTextMatching, ExactTextMatching from pyrit.memory.central_memory import CentralMemory @@ -27,7 +27,7 @@ async def test_score_async_unsupported_data_type_returns_false( scorer = SubStringScorer(substring="test", categories=["new_category"]) # With raise_on_no_valid_pieces=False (default), returns False for unsupported data types - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 assert scores[0].get_value() is False assert "No supported pieces" in scores[0].score_rationale diff --git a/tests/unit/score/test_true_false_composite_scorer.py b/tests/unit/score/test_true_false_composite_scorer.py index 3656fe5a9c..267b9e325d 100644 --- a/tests/unit/score/test_true_false_composite_scorer.py +++ b/tests/unit/score/test_true_false_composite_scorer.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock import pytest +from unit.mocks import store_message from pyrit.memory.central_memory import CentralMemory from pyrit.models import ComponentIdentifier, MessagePiece, Score, ScoringExpectation @@ -83,7 +84,7 @@ def false_scorer(patch_central_database): async def test_composite_scorer_and_all_true(mock_request, true_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer, true_scorer]) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -93,7 +94,7 @@ async def test_composite_scorer_and_all_true(mock_request, true_scorer): async def test_composite_scorer_and_one_false(mock_request, true_scorer, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer, false_scorer]) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a false score" in scores[0].score_rationale @@ -103,7 +104,7 @@ async def test_composite_scorer_and_one_false(mock_request, true_scorer, false_s async def test_composite_scorer_or_all_false(mock_request, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.OR, scorers=[false_scorer, false_scorer]) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a false score" in scores[0].score_rationale @@ -113,7 +114,7 @@ async def test_composite_scorer_or_all_false(mock_request, false_scorer): async def test_composite_scorer_or_one_true(mock_request, true_scorer, false_scorer): scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.OR, scorers=[true_scorer, false_scorer]) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -124,7 +125,7 @@ async def test_composite_scorer_majority_true(mock_request, true_scorer, false_s aggregator=TrueFalseScoreAggregator.MAJORITY, scorers=[true_scorer, true_scorer, false_scorer] ) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is True assert "This is a true score" in scores[0].score_rationale @@ -139,7 +140,7 @@ async def test_composite_scorer_majority_false(mock_request, true_scorer, false_ aggregator=TrueFalseScoreAggregator.MAJORITY, scorers=[true_scorer, false_scorer, false_scorer] ) - scores = await scorer.score_async(scorable=MessageScorable(message=mock_request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(mock_request))) assert len(scores) == 1 assert scores[0].get_value() is False assert "This is a true score" in scores[0].score_rationale @@ -166,7 +167,8 @@ async def test_composite_scorer_with_task(mock_request, true_scorer): task = "test task" scores = await scorer.score_async( - scorable=MessageScorable(message=mock_request), expectation=ScoringExpectation(objective=task) + scorable=MessageScorable.from_message(store_message(mock_request)), + expectation=ScoringExpectation(objective=task), ) assert len(scores) == 1 assert scores[0].objective == task @@ -178,17 +180,14 @@ def test_composite_scorer_empty_scorers_list(): TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[]) -async def test_composite_scorer_raises_when_message_piece_id_is_none(true_scorer, patch_central_database): - """Test that _score_async raises ValueError when message piece has no ID.""" +async def test_composite_scorer_anchors_where_its_children_anchored(true_scorer, patch_central_database): + """The aggregate is about whatever its children were about, so it anchors where they did.""" scorer = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[true_scorer]) + message = store_message(MessagePiece(role="user", original_value="test content").to_message()) - # Create a message with a piece whose id is None - piece = MessagePiece(role="user", original_value="test content") - piece.id = None - message = piece.to_message() + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) - with pytest.raises(RuntimeError, match="Message piece must have an ID"): - await scorer.score_async(scorable=MessageScorable(message=message)) + assert str(scores[0].message_piece_id) == str(message.get_piece().id) def test_get_chat_target_returns_first_available(patch_central_database): diff --git a/tests/unit/score/test_true_false_inverter.py b/tests/unit/score/test_true_false_inverter.py index ecb0bd924d..50003366c1 100644 --- a/tests/unit/score/test_true_false_inverter.py +++ b/tests/unit/score/test_true_false_inverter.py @@ -5,7 +5,7 @@ from unittest.mock import MagicMock, patch import pytest -from unit.mocks import get_image_message_piece +from unit.mocks import get_image_message_piece, store_message from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface @@ -28,7 +28,7 @@ async def test_score_async_unsupported_data_type_inverts_false_to_true( # With raise_on_no_valid_pieces=False (default), the inner scorer returns False, # and the inverter inverts it to True - scores = await scorer.score_async(scorable=MessageScorable(message=request)) + scores = await scorer.score_async(scorable=MessageScorable.from_message(store_message(request))) assert len(scores) == 1 # Inverter inverts False -> True assert scores[0].get_value() is True From d5f5245d344deb8c0afbae6afdc1dc4d392fa6d1 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 13:17:07 -0700 Subject: [PATCH 06/11] Describe a message as a scorable in one place The choice between loose content and a message reference was repeated in the legacy shim and in each of the three true/false wrappers. Move it to Scorer._scorable_for_message, which is the single bridge for the two callers that still hold a Message rather than the scorable it came from. The gist puts wrapper forwarding and the composite anchor fix in phase 2, and phase 1 changes no scorer body, so the wrappers keep resolving a message for now. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/score/scorer.py | 59 +++++++++++++------ .../float_scale_threshold_scorer.py | 11 +--- .../true_false/true_false_composite_scorer.py | 10 +--- .../true_false/true_false_inverter_scorer.py | 11 +--- 4 files changed, 45 insertions(+), 46 deletions(-) diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 800e8139a7..f4f4b75827 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -315,29 +315,54 @@ def _consolidate_legacy_inputs( ) if scorable is None: - # The caller asked about this message, not about everything stored alongside it, - # so the shim maps to the exact message rather than widening to its conversation. - # A message the caller built by hand has nothing in memory to name, so a single - # unpersisted piece becomes loose content. Both arms go away with the parameter. - legacy_message = cast("Message", message) - pieces = legacy_message.message_pieces - if len(pieces) == 1 and pieces[0].not_in_memory: - scorable = ContentScorable( - value=pieces[0].original_value, - data_type=pieces[0].original_value_data_type, - ) - else: - scorable = MessageScorable.from_message( - legacy_message, - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) + scorable = self._scorable_for_message( + cast("Message", message), + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) if objective is not None: expectation = ScoringExpectation(objective=objective) return scorable, expectation + @staticmethod + def _scorable_for_message( + message: Message, + *, + role_filter: ChatMessageRole | None = None, + skip_on_error_result: bool = False, + ) -> Scorable: + """ + Describe an in-hand message as the scorable that names it. + + The caller asked about this message, not about everything stored alongside it, so this + maps to the exact message rather than widening to its conversation. A message built by + hand has nothing in memory to name, so a single unpersisted piece becomes loose content. + + This bridges the two places that still hold a ``Message`` rather than the scorable it + came from: the deprecated ``message`` parameter, and the true/false wrappers that + forward to an inner scorer. Both go away in phase 2, and this goes with them. + + Args: + message (Message): The message to describe. + role_filter (ChatMessageRole | None): Role the piece must have to be scored. + skip_on_error_result (bool): Whether to skip the message when it holds an error. + + Returns: + Scorable: A ``ContentScorable`` for a single unpersisted piece, else a + ``MessageScorable`` naming the pieces. + """ + pieces = message.message_pieces + if len(pieces) == 1 and pieces[0].not_in_memory: + return ContentScorable.from_message(message) + + return MessageScorable.from_message( + message, + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) + @abstractmethod async def _score_scorable_async( self, diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index 26184a64a8..0b98b9a01f 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -10,7 +10,6 @@ from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleAggregatorFunc, FloatScaleScoreAggregator from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer -from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.score_utils import ORIGINAL_FLOAT_VALUE_KEY from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -99,16 +98,8 @@ async def _score_async( Returns: list[Score]: A list containing a single true/false Score object based on the threshold comparison. """ - # A message the caller built by hand has no ids to name, so it is content rather - # than a reference. Both arms go away when loose content is persisted in its own right. - pieces = message.message_pieces - scorable = ( - ContentScorable.from_message(message) - if len(pieces) == 1 and pieces[0].not_in_memory - else MessageScorable.from_message(message, role_filter=role_filter) - ) scores = await self._scorer.score_async( - scorable=scorable, + scorable=self._scorable_for_message(message, role_filter=role_filter), expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index 0c39791d59..ce1df70445 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -8,7 +8,6 @@ from pyrit.prompt_target import PromptTarget from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation -from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -98,14 +97,7 @@ async def _score_async( ValueError: If any constituent scorer does not return exactly one score. ValueError: If no scores are generated from the request response pieces. """ - # A message the caller built by hand has no ids to name, so it is content rather - # than a reference. Both arms go away when loose content is persisted in its own right. - pieces = message.message_pieces - scorable = ( - ContentScorable.from_message(message) - if len(pieces) == 1 and pieces[0].not_in_memory - else MessageScorable.from_message(message, role_filter=role_filter) - ) + scorable = self._scorable_for_message(message, role_filter=role_filter) tasks = [ scorer.score_async( scorable=scorable, diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index bde6bb318a..6d2480598d 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -8,7 +8,6 @@ from pyrit.prompt_target import PromptTarget from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation -from pyrit.score.scorable import ContentScorable, MessageScorable from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -74,16 +73,8 @@ async def _score_async( Returns: list[Score]: A list containing a single Score object with the inverted true/false value. """ - # A message the caller built by hand has no ids to name, so it is content rather - # than a reference. Both arms go away when loose content is persisted in its own right. - pieces = message.message_pieces - scorable = ( - ContentScorable.from_message(message) - if len(pieces) == 1 and pieces[0].not_in_memory - else MessageScorable.from_message(message, role_filter=role_filter) - ) scores = await self._scorer.score_async( - scorable=scorable, + scorable=self._scorable_for_message(message, role_filter=role_filter), expectation=ScoringExpectation(objective=objective) if objective is not None else None, ) inv_score = scores[0] From b8c92c2adb3b54fb287cb78ca35064355f2fa4dd Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 14:22:28 -0700 Subject: [PATCH 07/11] Always pass an expectation rather than a conditional one ScoringExpectation.objective is already optional, and a scorer reads it as expectation.objective if expectation else None, so an expectation holding no objective and no expectation at all mean the same thing. The ternary said otherwise in eight places and was dropped in a ninth, which made the callers look like they differed. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: d6f9fd6d-38b7-4841-bb94-c608033c1c04 --- pyrit/score/audio_transcript_scorer.py | 2 +- pyrit/score/scorer.py | 10 +++++----- pyrit/score/true_false/float_scale_threshold_scorer.py | 2 +- pyrit/score/true_false/true_false_composite_scorer.py | 9 ++------- pyrit/score/true_false/true_false_inverter_scorer.py | 2 +- 5 files changed, 10 insertions(+), 15 deletions(-) diff --git a/pyrit/score/audio_transcript_scorer.py b/pyrit/score/audio_transcript_scorer.py index aaf82af688..573a1f497c 100644 --- a/pyrit/score/audio_transcript_scorer.py +++ b/pyrit/score/audio_transcript_scorer.py @@ -188,7 +188,7 @@ async def _score_audio_async(self, *, message_piece: MessagePiece, objective: st # Score the transcript transcript_scores = await self.text_scorer.score_async( scorable=MessageScorable.from_message(text_message), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) # Add context to indicate this was scored from audio transcription diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index f4f4b75827..af3eb1aa08 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -640,7 +640,7 @@ async def score_text_async(self, text: str, *, objective: str | None = None) -> """ return await self.score_async( scorable=ContentScorable(value=text), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) async def score_image_async(self, image_path: str, *, objective: str | None = None) -> list[Score]: @@ -656,7 +656,7 @@ async def score_image_async(self, image_path: str, *, objective: str | None = No """ return await self.score_async( scorable=ContentScorable(value=image_path, data_type="image_path"), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) async def score_prompts_batch_async( @@ -852,7 +852,7 @@ async def score_response_async( scorable=MessageScorable.from_message( response, role_filter=role_filter, skip_on_error_result=skip_on_error_result ), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) aux_scores, obj_scores = await asyncio.gather(aux_task, obj_task) result["auxiliary_scores"] = aux_scores @@ -862,7 +862,7 @@ async def score_response_async( scorable=MessageScorable.from_message( response, role_filter=role_filter, skip_on_error_result=skip_on_error_result ), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) result["objective_scores"] = obj_scores return result @@ -900,7 +900,7 @@ async def score_response_multiple_scorers_async( scorable = MessageScorable.from_message( response, role_filter=role_filter, skip_on_error_result=skip_on_error_result ) - expectation = ScoringExpectation(objective=objective) if objective is not None else None + expectation = ScoringExpectation(objective=objective) tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in scorers] if not tasks: diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index 0b98b9a01f..04c2332b8b 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -100,7 +100,7 @@ async def _score_async( """ scores = await self._scorer.score_async( scorable=self._scorable_for_message(message, role_filter=role_filter), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) # Aggregator handles 0-many scores and returns exactly one result (or raises if configured) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index ce1df70445..88538467cc 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -98,13 +98,8 @@ async def _score_async( ValueError: If no scores are generated from the request response pieces. """ scorable = self._scorable_for_message(message, role_filter=role_filter) - tasks = [ - scorer.score_async( - scorable=scorable, - expectation=ScoringExpectation(objective=objective) if objective is not None else None, - ) - for scorer in self._scorers - ] + expectation = ScoringExpectation(objective=objective) + tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in self._scorers] # Run all response scorings concurrently score_list_results = await asyncio.gather(*tasks) diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index 6d2480598d..7184642fa9 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -75,7 +75,7 @@ async def _score_async( """ scores = await self._scorer.score_async( scorable=self._scorable_for_message(message, role_filter=role_filter), - expectation=ScoringExpectation(objective=objective) if objective is not None else None, + expectation=ScoringExpectation(objective=objective), ) inv_score = scores[0] From fb4894a123a860612f5a7018496235f8c70a7d8d Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Fri, 14 Aug 2026 16:35:25 -0700 Subject: [PATCH 08/11] Refine scorable ownership and resolution Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 1e5bf78a-8e94-4839-b6bf-7f29a6a19a38 --- pyrit/executor/attack/multi_turn/crescendo.py | 2 + .../executor/attack/multi_turn/red_teaming.py | 4 +- pyrit/models/__init__.py | 13 +- pyrit/models/score/__init__.py | 8 +- pyrit/models/score/scorable.py | 82 ++++ pyrit/score/audio_transcript_scorer.py | 3 +- pyrit/score/conversation_scorer.py | 2 +- pyrit/score/float_scale/float_scale_scorer.py | 16 +- pyrit/score/message_scorable_resolver.py | 80 ++++ pyrit/score/message_scorer.py | 323 ++++++++++++- pyrit/score/scorable.py | 170 +------ pyrit/score/scorer.py | 427 ++++-------------- .../float_scale_threshold_scorer.py | 17 +- .../true_false/true_false_composite_scorer.py | 17 +- .../true_false/true_false_inverter_scorer.py | 17 +- pyrit/score/true_false/true_false_scorer.py | 9 +- .../attack/multi_turn/test_crescendo.py | 6 +- .../multi_turn/test_crescendo_resilience.py | 3 +- tests/unit/models/test_scorable.py | 86 ++++ .../score/test_message_scorable_resolver.py | 95 ++++ tests/unit/score/test_message_scorer.py | 92 +++- tests/unit/score/test_scorable.py | 175 ------- tests/unit/score/test_scorer.py | 58 ++- 23 files changed, 954 insertions(+), 751 deletions(-) create mode 100644 pyrit/models/score/scorable.py create mode 100644 pyrit/score/message_scorable_resolver.py create mode 100644 tests/unit/models/test_scorable.py create mode 100644 tests/unit/score/test_message_scorable_resolver.py delete mode 100644 tests/unit/score/test_scorable.py diff --git a/pyrit/executor/attack/multi_turn/crescendo.py b/pyrit/executor/attack/multi_turn/crescendo.py index 05fad46acd..0d9ce56cb7 100644 --- a/pyrit/executor/attack/multi_turn/crescendo.py +++ b/pyrit/executor/attack/multi_turn/crescendo.py @@ -45,6 +45,7 @@ SelfAskRefusalScorer, SelfAskScaleScorer, ) +from pyrit.score.message_scorer import MessageScoringOptions from pyrit.score.score_utils import normalize_score_to_float if TYPE_CHECKING: @@ -674,6 +675,7 @@ async def _check_refusal_async(self, context: CrescendoAttackContext, objective: scores = await self._refusal_scorer.score_async( scorable=MessageScorable.from_message(context.last_response), expectation=ScoringExpectation(objective=objective), + message_options=MessageScoringOptions(skip_on_error_result=False), ) return scores[0] diff --git a/pyrit/executor/attack/multi_turn/red_teaming.py b/pyrit/executor/attack/multi_turn/red_teaming.py index 7216254cdb..06af753d31 100644 --- a/pyrit/executor/attack/multi_turn/red_teaming.py +++ b/pyrit/executor/attack/multi_turn/red_teaming.py @@ -40,6 +40,7 @@ from pyrit.prompt_target import CapabilityName from pyrit.prompt_target.common.target_requirements import TargetRequirements from pyrit.score import MessageScorable +from pyrit.score.message_scorer import MessageScoringOptions if TYPE_CHECKING: from collections.abc import Callable @@ -525,8 +526,9 @@ async def _score_response_async(self, *, context: MultiTurnAttackContext[Any]) - ): # score_async handles blocked, filtered, other errors scoring_results = await self._objective_scorer.score_async( - scorable=MessageScorable.from_message(context.last_response, role_filter="assistant"), + scorable=MessageScorable.from_message(context.last_response), expectation=ScoringExpectation(objective=context.objective), + message_options=MessageScoringOptions(role_filter="assistant"), ) objective_scores = scoring_results diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 2cbbcac1c2..ca03c5c7e1 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -78,7 +78,15 @@ from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT from pyrit.models.retry_event import RetryEvent -from pyrit.models.score import Score, ScoreType, ScoringExpectation, UnvalidatedScore +from pyrit.models.score import ( + ContentScorable, + MessageScorable, + Scorable, + Score, + ScoreType, + ScoringExpectation, + UnvalidatedScore, +) # Seeds - import from new seeds submodule for forward compatibility # Also keep imports from old locations for backward compatibility @@ -140,6 +148,7 @@ "ConversationRetryReason", "ConversationStats", "ConversationType", + "ContentScorable", "construct_response_from_request", "display_choices", "EmbeddingData", @@ -170,6 +179,7 @@ "MEDIA_PATH_DATA_TYPES", "Message", "MessagePiece", + "MessageScorable", "Modality", "NextMessageSystemPromptPaths", "ObjectiveTargetEvaluationIdentifier", @@ -183,6 +193,7 @@ "QuestionChoice", "REGISTRY_NAME_PATTERN", "ScaleDescription", + "Scorable", "Score", "ScoreType", "ScoringExpectation", diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py index bc357c13c1..28151b0d4b 100644 --- a/pyrit/models/score/__init__.py +++ b/pyrit/models/score/__init__.py @@ -2,18 +2,22 @@ # Licensed under the MIT license. """ -Score types: what a scorer is scored against, and the result. +Score types: what a scorer looks at, what it scores against, and the result. A scorer takes two inputs — a ``Scorable`` (what to look at) and a ``ScoringExpectation`` (what to look for) — and returns ``Score`` objects. Scorables -resolve themselves against memory, so they live in ``pyrit.score`` rather than here. +are inert canonical data; scoring-layer resolvers acquire the evidence they name. """ from pyrit.models.score.expectation import ScoringExpectation +from pyrit.models.score.scorable import ContentScorable, MessageScorable, Scorable from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore __all__ = [ "ComponentIdentifierField", + "ContentScorable", + "MessageScorable", + "Scorable", "Score", "ScoreType", "ScoringExpectation", diff --git a/pyrit/models/score/scorable.py b/pyrit/models/score/scorable.py new file mode 100644 index 0000000000..9185a3179a --- /dev/null +++ b/pyrit/models/score/scorable.py @@ -0,0 +1,82 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import uuid # noqa: TC003 (runtime-required by dataclass field annotations) +from abc import ABC +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from pyrit.models.literals import PromptDataType # noqa: TC001 (runtime-required by dataclass field annotations) + +if TYPE_CHECKING: + from pyrit.models.messages.message import Message + + +class Scorable(ABC): # noqa: B024 root type; each scorer family declares its own contract + """ + What a scorer looks at. + + A scorable is normally an inert reference: it names the evidence instead of carrying + or acquiring it. ``ContentScorable`` is the exception, because loose content has + nothing behind it to point at. A scorer-family resolver acquires the named evidence. + """ + + +@dataclass(frozen=True, kw_only=True) +class MessageScorable(Scorable): + """ + Specific message pieces, named by id. + + This names one message, or a subset of its pieces. Loose content that was never + persisted has no ids to name, so it is a ``ContentScorable`` instead. + """ + + message_piece_ids: tuple[uuid.UUID | str, ...] + + @classmethod + def from_message( + cls, + message: Message, + ) -> MessageScorable: + """ + Name the pieces of a persisted message. + + Args: + message (Message): The message whose pieces to name. + + Returns: + MessageScorable: A scorable naming the message's pieces. + """ + return cls(message_piece_ids=tuple(piece.id for piece in message.message_pieces)) + + +@dataclass(frozen=True, kw_only=True) +class ContentScorable(Scorable): + """ + Loose content with no conversation behind it. + + This names content, not a message: there is no role or error state. A message-family + resolver adapts it for existing message scorers. + """ + + value: str + data_type: PromptDataType = "text" + + @classmethod + def from_message(cls, message: Message) -> ContentScorable: + """ + Describe the converted content of a single-piece ephemeral message. + + Scorers consume ``converted_value``, so this adapter preserves the converted value + and data type rather than the pre-conversion input. + + Args: + message (Message): The ephemeral message whose converted content to take. + + Returns: + ContentScorable: A scorable holding the converted message content. + """ + piece = message.get_piece() + return cls(value=piece.converted_value, data_type=piece.converted_value_data_type) diff --git a/pyrit/score/audio_transcript_scorer.py b/pyrit/score/audio_transcript_scorer.py index 573a1f497c..b08e64d2c0 100644 --- a/pyrit/score/audio_transcript_scorer.py +++ b/pyrit/score/audio_transcript_scorer.py @@ -11,8 +11,7 @@ from pyrit.converter import AzureSpeechAudioToTextConverter from pyrit.memory import CentralMemory -from pyrit.models import MessagePiece, Score, ScoringExpectation -from pyrit.score.scorable import MessageScorable +from pyrit.models import MessagePiece, MessageScorable, Score, ScoringExpectation from pyrit.score.scorer import Scorer logger = logging.getLogger(__name__) diff --git a/pyrit/score/conversation_scorer.py b/pyrit/score/conversation_scorer.py index c3cb68230f..13f4274773 100644 --- a/pyrit/score/conversation_scorer.py +++ b/pyrit/score/conversation_scorer.py @@ -209,7 +209,7 @@ class DynamicConversationScorer(ConversationScorer, scorer_base_class): # type: def __init__(self) -> None: # Initialize with the validator and wrapped scorer - Scorer.__init__(self, validator=validator or ConversationScorer._DEFAULT_VALIDATOR) + MessageScorer.__init__(self, validator=validator or ConversationScorer._DEFAULT_VALIDATOR) self._wrapped_scorer = wrapped_scorer def _get_wrapped_scorer(self) -> MessageScorer: diff --git a/pyrit/score/float_scale/float_scale_scorer.py b/pyrit/score/float_scale/float_scale_scorer.py index 6206dadeea..02413eaa20 100644 --- a/pyrit/score/float_scale/float_scale_scorer.py +++ b/pyrit/score/float_scale/float_scale_scorer.py @@ -10,6 +10,7 @@ if TYPE_CHECKING: from pyrit.prompt_target.common.prompt_target import PromptTarget + from pyrit.score.message_scorable_resolver import MessageScorableResolver from pyrit.score.scorer_evaluation.scorer_metrics import HarmScorerMetrics from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -35,7 +36,13 @@ class FloatScaleScorer(MessageScorer): "blocked = True") should override ``_score_piece_async`` or ``_build_fallback_score``. """ - def __init__(self, *, validator: ScorerPromptValidator, chat_target: PromptTarget | None = None) -> None: + def __init__( + self, + *, + validator: ScorerPromptValidator, + chat_target: PromptTarget | None = None, + message_resolver: MessageScorableResolver | None = None, + ) -> None: """ Initialize the FloatScaleScorer. @@ -43,8 +50,13 @@ def __init__(self, *, validator: ScorerPromptValidator, chat_target: PromptTarge validator: A validator object used to validate scores. chat_target: Optional chat target used by the scorer, forwarded to the base class for validation against ``TARGET_REQUIREMENTS``. + message_resolver: Message evidence resolver. """ - super().__init__(validator=validator, chat_target=chat_target) + super().__init__( + validator=validator, + chat_target=chat_target, + message_resolver=message_resolver, + ) def _build_fallback_score( self, *, message: Message, objective: str | None, scorer_response_blocked: bool = False diff --git a/pyrit/score/message_scorable_resolver.py b/pyrit/score/message_scorable_resolver.py new file mode 100644 index 0000000000..18582fc379 --- /dev/null +++ b/pyrit/score/message_scorable_resolver.py @@ -0,0 +1,80 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from pyrit.models import ( + ContentScorable, + Message, + MessagePiece, + MessageScorable, + Scorable, + group_message_pieces_into_conversations, +) + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +class MessageScorableResolver: + """Acquire message-shaped evidence for a ``MessageScorer``.""" + + def resolve(self, *, scorable: Scorable, memory: MemoryInterface) -> Message: + """ + Resolve supported scorables to the message view consumed by message scorers. + + Args: + scorable (Scorable): A message reference or loose content. + memory (MemoryInterface): Memory used to resolve message references. + + Returns: + Message: The message view to score. + + Raises: + TypeError: If the scorable is not message-shaped. + ValueError: If referenced pieces are missing or do not form one message. + """ + if isinstance(scorable, MessageScorable): + return self._resolve_message_reference(scorable=scorable, memory=memory) + if isinstance(scorable, ContentScorable): + return self._adapt_content(scorable=scorable) + raise TypeError( + f"Message scorers cannot score {type(scorable).__name__}. Pass a MessageScorable or a ContentScorable." + ) + + @staticmethod + def _resolve_message_reference(*, scorable: MessageScorable, memory: MemoryInterface) -> Message: + pieces = memory.get_message_pieces(prompt_ids=list(scorable.message_piece_ids)) + wanted = {str(piece_id) for piece_id in scorable.message_piece_ids} + pieces = [piece for piece in pieces if str(piece.id) in wanted] + found = {str(piece.id) for piece in pieces} + missing = [str(piece_id) for piece_id in scorable.message_piece_ids if str(piece_id) not in found] + if missing: + raise ValueError(f"No message pieces found in memory for ids {missing}.") + + conversations = group_message_pieces_into_conversations(pieces) + messages = [message for conversation in conversations for message in conversation] + if len(messages) != 1: + raise ValueError( + f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " + "Reference pieces from a single message." + ) + + resolved = messages[0] + by_id = {str(piece.id): piece for piece in resolved.message_pieces} + resolved.message_pieces = [by_id[str(piece_id)] for piece_id in scorable.message_piece_ids] + return resolved + + @staticmethod + def _adapt_content(*, scorable: ContentScorable) -> Message: + piece = MessagePiece( + role="user", + original_value=scorable.value, + converted_value=scorable.value, + original_value_data_type=scorable.data_type, + converted_value_data_type=scorable.data_type, + ) + piece.not_in_memory = True + return piece.to_message() diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index 5c6dd4dc86..9e454d8d7b 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -6,19 +6,41 @@ import asyncio import logging from abc import abstractmethod -from typing import TYPE_CHECKING +from dataclasses import dataclass +from typing import TYPE_CHECKING, cast +from pyrit.common.deprecation import print_deprecation_message from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException -from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable -from pyrit.score.scorer import Scorer +from pyrit.models import ( + ChatMessageRole, + ContentScorable, + Message, + MessagePiece, + MessageScorable, + PromptResponseError, + Scorable, + Score, + ScoringExpectation, +) +from pyrit.score.message_scorable_resolver import MessageScorableResolver +from pyrit.score.scorer import LEGACY_SCORE_ASYNC_REMOVED_IN, Scorer if TYPE_CHECKING: from pyrit.memory import MemoryInterface - from pyrit.models import Message, MessagePiece, Score, ScoringExpectation + from pyrit.prompt_target import PromptTarget + from pyrit.score.scorer_prompt_validator import ScorerPromptValidator logger = logging.getLogger(__name__) +@dataclass(frozen=True, kw_only=True) +class MessageScoringOptions: + """Message-only scoring policy that is not part of evidence identity.""" + + role_filter: ChatMessageRole | None = None + skip_on_error_result: bool = False + + def extract_objective_from_previous_turn(*, message: Message, memory: MemoryInterface) -> str: """ Read the text of the turn before an assistant message and use it as the objective. @@ -67,17 +89,163 @@ class MessageScorer(Scorer): Every message-shaped concern lives here: substituting refusal and blocked content, validating pieces, applying the role and error filters, and falling back to a neutral score. ``Scorer`` stays agnostic about what a scorable is, so scorers over other kinds of - evidence can sit beside this one. The scorable resolves itself to a ``Message``. + evidence can sit beside this one. A ``MessageScorableResolver`` acquires the message; + the scorable remains inert. Subclasses implement ``_score_async``, which still receives a ``Message``. """ + #: When True, blocked responses that contain partial content are scored using that + #: content instead of being filtered out or short-circuited. + score_blocked_content: bool = False + + #: When False, a blocked response from the scorer's own LLM produces the scorer + #: family's neutral fallback score instead of raising. + raise_if_scorer_blocks: bool = True + + def __init__( + self, + *, + validator: ScorerPromptValidator, + chat_target: PromptTarget | None = None, + message_resolver: MessageScorableResolver | None = None, + ) -> None: + """ + Initialize message-specific scoring dependencies. + + Args: + validator (ScorerPromptValidator): Validator for message pieces. + chat_target (PromptTarget | None): Optional target used by the scorer. + message_resolver (MessageScorableResolver | None): Evidence resolver. + """ + self._validator = validator + self._message_resolver = message_resolver or MessageScorableResolver() + super().__init__(chat_target=chat_target) + + async def score_async( + self, + message: Message | None = None, + *, + scorable: Scorable | None = None, + expectation: ScoringExpectation | None = None, + message_options: MessageScoringOptions | None = None, + objective: str | None = None, + role_filter: ChatMessageRole | None = None, + skip_on_error_result: bool | None = None, + infer_objective_from_request: bool | None = None, + ) -> list[Score]: + """ + Score message-shaped evidence, including the deprecated message API. + + Args: + message (Message | None): Deprecated in-hand message. + scorable (Scorable | None): Message-shaped evidence to acquire. + expectation (ScoringExpectation | None): What to look for. + message_options (MessageScoringOptions | None): Message-family policy. + objective (str | None): Deprecated objective string. + role_filter (ChatMessageRole | None): Deprecated role policy. + skip_on_error_result (bool | None): Deprecated error policy. ``None`` means omitted. + infer_objective_from_request (bool | None): Deprecated inference policy. + + Returns: + list[Score]: The persisted scores, or an empty list when policy skips the message. + """ + resolved_scorable, resolved_expectation, options, infer_objective = self._consolidate_message_inputs( + message=message, + scorable=scorable, + expectation=expectation, + message_options=message_options, + objective=objective, + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + infer_objective_from_request=infer_objective_from_request, + ) + scores = await self._score_message_scorable_async( + scorable=resolved_scorable, + expectation=resolved_expectation, + options=options, + infer_objective_from_request=infer_objective, + ) + return self._validate_and_persist_scores(scores=scores) + + def _consolidate_message_inputs( + self, + *, + message: Message | None, + scorable: Scorable | None, + expectation: ScoringExpectation | None, + message_options: MessageScoringOptions | None, + objective: str | None, + role_filter: ChatMessageRole | None, + skip_on_error_result: bool | None, + infer_objective_from_request: bool | None, + ) -> tuple[Scorable, ScoringExpectation | None, MessageScoringOptions, bool]: + if message is not None and scorable is not None: + raise ValueError("Pass either 'message' or 'scorable', not both.") + if message is None and scorable is None: + raise ValueError("Either 'message' or 'scorable' must be provided.") + if objective is not None and expectation is not None: + raise ValueError("Pass either 'objective' or 'expectation', not both.") + if message_options is not None and (role_filter is not None or skip_on_error_result is not None): + raise ValueError("Pass either 'message_options' or legacy message policy arguments, not both.") + + uses_legacy_parameters = ( + message is not None + or objective is not None + or role_filter is not None + or skip_on_error_result is not None + or infer_objective_from_request is not None + ) + if uses_legacy_parameters: + print_deprecation_message( + old_item="Scorer.score_async(message=..., objective=..., role_filter=..., " + "skip_on_error_result=..., infer_objective_from_request=...)", + new_item="Scorer.score_async(scorable=..., expectation=..., message_options=...)", + removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, + ) + + if scorable is not None: + resolved_scorable = scorable + else: + legacy_message = cast("Message", message) + resolved_scorable = ( + ContentScorable.from_message(legacy_message) + if len(legacy_message.message_pieces) == 1 and legacy_message.get_piece().not_in_memory + else MessageScorable.from_message(legacy_message) + ) + resolved_expectation = ScoringExpectation(objective=objective) if objective is not None else expectation + options = message_options or MessageScoringOptions( + role_filter=role_filter, + skip_on_error_result=skip_on_error_result or False, + ) + return resolved_scorable, resolved_expectation, options, bool(infer_objective_from_request) + async def _score_scorable_async( self, *, scorable: Scorable, expectation: ScoringExpectation | None, - infer_objective_from_request: bool = False, + ) -> list[Score]: + """ + Score message-shaped evidence with default message policy. + + Returns: + list[Score]: The scores produced from the resolved message. + """ + return await self._score_message_scorable_async( + scorable=scorable, + expectation=expectation, + options=MessageScoringOptions(), + infer_objective_from_request=False, + ) + + async def _score_message_scorable_async( + self, + *, + scorable: Scorable, + expectation: ScoringExpectation | None, + options: MessageScoringOptions, + infer_objective_from_request: bool, ) -> list[Score]: """ Resolve a message scorable and score the message it names. @@ -85,6 +253,7 @@ async def _score_scorable_async( Args: scorable (Scorable): A ``MessageScorable`` or a ``ContentScorable``. expectation (ScoringExpectation | None): What to look for. + options (MessageScoringOptions): Message-only scoring policy. infer_objective_from_request (bool): Deprecated; read the objective from the previous turn when the expectation carries none. @@ -98,21 +267,7 @@ async def _score_scorable_async( PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). """ - if isinstance(scorable, MessageScorable): - message = scorable.resolve_message(memory=self._memory) - role_filter = scorable.role_filter - skip_on_error_result = scorable.skip_on_error_result - elif isinstance(scorable, ContentScorable): - # Loose content names no message, so there is nothing to filter on. Phase 2 - # persists the content and this arm goes away. - message = scorable.to_ephemeral_message() - role_filter = None - skip_on_error_result = False - else: - raise TypeError( - f"{self.__class__.__name__} scores messages, so it cannot score {type(scorable).__name__}. " - "Pass a MessageScorable or a ContentScorable." - ) + message = self._message_resolver.resolve(scorable=scorable, memory=self._memory) objective = expectation.objective if expectation else None @@ -128,11 +283,11 @@ async def _score_scorable_async( self._validator.validate(scoring_message, objective=objective) - if role_filter is not None and message.get_piece().role != role_filter: + if options.role_filter is not None and message.get_piece().role != options.role_filter: logger.debug("Skipping scoring due to role filter mismatch.") return [] - if skip_on_error_result and self._should_skip_on_error(message): + if options.skip_on_error_result and self._should_skip_on_error(message): return [] if infer_objective_from_request and (not objective): @@ -268,3 +423,125 @@ def _drop_ephemeral_score_links(*, message: Message, scores: list[Score]) -> Non for score in scores: if score.message_piece_id in ephemeral_piece_ids: score.message_piece_id = None # type: ignore[ty:invalid-assignment] + + @staticmethod + def _create_scoring_text_piece( + *, + piece: MessagePiece, + content: str, + response_error: PromptResponseError, + ) -> MessagePiece: + """ + Create a text scoring view that retains the persisted piece identity. + + Returns: + MessagePiece: The text scoring view. + """ + return MessagePiece( + id=piece.id, + role=piece.api_role, + original_value=piece.original_value, + converted_value=content, + original_value_data_type=piece.original_value_data_type, + converted_value_data_type="text", + conversation_id=piece.conversation_id, + sequence=piece.sequence, + prompt_metadata=dict(piece.prompt_metadata), + converter_identifiers=list(piece.converter_identifiers), # type: ignore[arg-type] + response_error=response_error, + timestamp=piece.timestamp, + original_prompt_id=piece.original_prompt_id, + not_in_memory=piece.not_in_memory, + ) + + @classmethod + def _create_text_piece_from_blocked(cls, piece: MessagePiece) -> MessagePiece | None: + """ + Create a text scoring view from a blocked piece's partial content. + + Returns: + MessagePiece | None: The scoring view, or None when content is unavailable. + """ + partial_content = str(piece.prompt_metadata.get("partial_content", "")) + if not partial_content: + return None + return cls._create_scoring_text_piece( + piece=piece, + content=partial_content, + response_error="none", + ) + + @classmethod + def _create_text_piece_from_structured_refusal(cls, piece: MessagePiece) -> MessagePiece | None: + """ + Create a blocked text scoring view for an SDK-provided refusal. + + Returns: + MessagePiece | None: The scoring view, or None when there is no refusal. + """ + refusal = piece.structured_refusal + if not refusal: + return None + return cls._create_scoring_text_piece( + piece=piece, + content=refusal, + response_error="blocked", + ) + + def _apply_structured_refusal_substitution(self, message: Message) -> Message: + """ + Expose structured refusal explanations while preserving blocked semantics. + + Returns: + Message: The substituted message, or the original message. + """ + substituted = False + new_pieces: list[MessagePiece] = [] + for piece in message.message_pieces: + substitute = self._create_text_piece_from_structured_refusal(piece) + if substitute: + new_pieces.append(substitute) + substituted = True + continue + new_pieces.append(piece) + return Message(message_pieces=new_pieces) if substituted else message + + def _apply_blocked_content_substitution(self, message: Message) -> Message: + """ + Replace blocked pieces that have partial content with text scoring views. + + Returns: + Message: The substituted message, or the original message. + """ + substituted = False + new_pieces: list[MessagePiece] = [] + for piece in message.message_pieces: + if piece.is_blocked() and "partial_content" in piece.prompt_metadata: + substitute = self._create_text_piece_from_blocked(piece) + if substitute: + new_pieces.append(substitute) + substituted = True + continue + new_pieces.append(piece) + return Message(message_pieces=new_pieces) if substituted else message + + @abstractmethod + def _build_fallback_score( + self, + *, + message: Message, + objective: str | None, + scorer_response_blocked: bool = False, + ) -> list[Score]: + """ + Return the scorer family's neutral result when message evidence is unscoreable. + + Args: + message (Message): The message-shaped evidence. + objective (str | None): The objective associated with this call. + scorer_response_blocked (bool): Whether the scorer's own LLM was blocked. + + Returns: + list[Score]: One or more fallback scores. + """ + ... diff --git a/pyrit/score/scorable.py b/pyrit/score/scorable.py index f47fd666b7..46d56b0c56 100644 --- a/pyrit/score/scorable.py +++ b/pyrit/score/scorable.py @@ -1,168 +1,12 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -from __future__ import annotations +"""Compatibility exports for scorables, canonically owned by ``pyrit.models.score``.""" -import uuid # noqa: TC003 (runtime-required by dataclass field annotations) -from abc import ABC -from dataclasses import dataclass -from typing import TYPE_CHECKING +from pyrit.models import ContentScorable, MessageScorable, Scorable -from pyrit.models import ( # noqa: TC001 (runtime-required by dataclass field annotations) - ChatMessageRole, - Message, - MessagePiece, - PromptDataType, - group_message_pieces_into_conversations, -) - -if TYPE_CHECKING: - from pyrit.memory import MemoryInterface - - -class Scorable(ABC): # noqa: B024 root type; each scorer family declares its own contract - """ - What a scorer looks at. - - A scorable is normally a reference: it names the evidence instead of carrying it, and - the scorer resolves it. ``ContentScorable`` is the exception, because loose content has - nothing behind it to point at. Each scorer accepts the scorable kinds it supports and - rejects the rest. - """ - - -@dataclass(frozen=True, kw_only=True) -class MessageScorable(Scorable): - """ - Specific message pieces, resolved from memory by id. - - This names one message, or a subset of its pieces. Loose content that was never - persisted has no ids to name, so it is a ``ContentScorable`` instead. - """ - - message_piece_ids: tuple[uuid.UUID | str, ...] - - # The design proposal puts role_filter on ConversationScorable only and leaves this - # open as its phase-1 decision, on the grounds that naming pieces is already a - # selection and filtering them again is a second one. Attacks do filter a named - # response by role today, so both fields stay here until ConversationScorable arrives - # and can own the selection instead. - role_filter: ChatMessageRole | None = None - skip_on_error_result: bool = False - - def resolve_message(self, *, memory: MemoryInterface) -> Message: - """ - Return the single message the referenced pieces form. - - Args: - memory (MemoryInterface): Memory holding the pieces. - - Returns: - Message: The message to score. - - Raises: - ValueError: If any referenced piece is not in memory, or the pieces do not form - exactly one message. - """ - pieces = memory.get_message_pieces(prompt_ids=list(self.message_piece_ids)) - wanted = {str(piece_id) for piece_id in self.message_piece_ids} - # Only the named pieces count, whatever else the filter happened to return. - pieces = [piece for piece in pieces if str(piece.id) in wanted] - found = {str(piece.id) for piece in pieces} - missing = [str(piece_id) for piece_id in self.message_piece_ids if str(piece_id) not in found] - if missing: - raise ValueError(f"No message pieces found in memory for ids {missing}.") - - conversations = group_message_pieces_into_conversations(pieces) - messages = [message for conversation in conversations for message in conversation] - if len(messages) != 1: - raise ValueError( - f"Expected the referenced pieces to form exactly one message, got {len(messages)}. " - "Reference pieces from a single message." - ) - - # Memory returns pieces in its own order, but the scorable names an ordered tuple, - # and a multi-piece message reads differently if its pieces are shuffled. - resolved = messages[0] - by_id = {str(piece.id): piece for piece in resolved.message_pieces} - resolved.message_pieces = [by_id[str(piece_id)] for piece_id in self.message_piece_ids] - return resolved - - @classmethod - def from_message( - cls, - message: Message, - *, - role_filter: ChatMessageRole | None = None, - skip_on_error_result: bool = False, - ) -> MessageScorable: - """ - Name the pieces of a persisted message. - - A scorable is a reference, so a message in hand is scored by naming its pieces - rather than by travelling through the call. Use ``scorable_from_message`` instead - when the message may not be persisted. - - Args: - message (Message): The message whose pieces to name. - role_filter (ChatMessageRole | None): Only score the message when it has this role. - skip_on_error_result (bool): Skip scoring when the message is an error result. - - Returns: - MessageScorable: A scorable naming the message's pieces. - """ - return cls( - message_piece_ids=tuple(piece.id for piece in message.message_pieces), - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) - - -@dataclass(frozen=True, kw_only=True) -class ContentScorable(Scorable): - """ - Loose content with no conversation behind it. - - This names content, not a message: there is no role and no error state, so the filters - a ``MessageScorable`` carries would be meaningless here. A message scorer adapts it - with ``to_ephemeral_message``. - """ - - value: str - data_type: PromptDataType = "text" - - def to_ephemeral_message(self) -> Message: - """ - Wrap the content as a message so a message scorer can read it. - - This is an adapter, not a resolution: no such message exists until this call builds - one, and it is marked as never persisted. Phase 2 stores loose content as a row of - its own and retires this. - - Returns: - Message: A throwaway message holding the content. - """ - piece = MessagePiece( - role="user", - original_value=self.value, - original_value_data_type=self.data_type, - ) - piece.not_in_memory = True - return Message(message_pieces=[piece]) - - @classmethod - def from_message(cls, message: Message) -> ContentScorable: - """ - Take the content out of a message that was never persisted. - - A message the caller built by hand has no ids to name, so it is content rather than - a reference. Only the first piece is taken, because loose content is a single value. - - Args: - message (Message): The unpersisted message whose content to take. - - Returns: - ContentScorable: A scorable holding the message's content. - """ - piece = message.get_piece() - return cls(value=piece.original_value, data_type=piece.original_value_data_type) +__all__ = [ + "ContentScorable", + "MessageScorable", + "Scorable", +] diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index af3eb1aa08..0654006f1b 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -14,10 +14,11 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, + ContentScorable, Identifiable, Message, - MessagePiece, - PromptResponseError, + MessageScorable, + Scorable, Score, ScorerEvaluationIdentifier, ScorerIdentifier, @@ -26,7 +27,6 @@ ) from pyrit.prompt_target.batch_helper import batch_task_async from pyrit.prompt_target.common.target_requirements import TargetRequirements -from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable if TYPE_CHECKING: from collections.abc import Sequence @@ -35,7 +35,6 @@ from pyrit.score.scorer_evaluation.metrics_type import RegistryUpdateBehavior from pyrit.score.scorer_evaluation.scorer_evaluator import ScorerEvalDatasetFiles from pyrit.score.scorer_evaluation.scorer_metrics import ScorerMetrics - from pyrit.score.scorer_prompt_validator import ScorerPromptValidator logger = logging.getLogger(__name__) @@ -65,24 +64,6 @@ class Scorer(Identifiable, abc.ABC): _identifier: ComponentIdentifier | None = None - #: When True, blocked responses that contain partial content - #: (in prompt_metadata["partial_content"]) will be scored using that content - #: instead of being filtered out or short-circuited. - #: Set this on scorer instances before use. Defaults to False. - #: - #: Note: Partial content extraction is supported for ``OpenAIChatTarget`` - #: (Chat Completions API) and ``OpenAIResponseTarget`` (Responses API). - score_blocked_content: bool = False - - #: Controls what happens when the scorer's *own* LLM response is blocked by content - #: filtering (common in red-teaming, since the scorer's rationale quotes harmful content). - #: When True (default), scoring raises ``ScorerLLMResponseBlockedException`` — a blocked - #: scorer endpoint is treated as a real error. When False, scoring returns the scorer's - #: type default instead (False for true/false scorers, 0.0 for float-scale). This is - #: distinct from ``score_blocked_content``, which concerns the target-under-test response. - #: Set this on scorer instances before use. Defaults to True. - raise_if_scorer_blocks: bool = True - def __init_subclass__(cls, **kwargs: Any) -> None: """ Enforce the keyword-only constructor contract on subclasses. @@ -95,16 +76,14 @@ def __init_subclass__(cls, **kwargs: Any) -> None: enforce_keyword_only_init(cls, base_name="Scorer") - def __init__(self, *, validator: ScorerPromptValidator, chat_target: PromptTarget | None = None) -> None: + def __init__(self, *, chat_target: PromptTarget | None = None) -> None: """ Initialize the Scorer. Args: - validator (ScorerPromptValidator): Validator for message pieces and scorer configuration. chat_target (PromptTarget | None): Chat target used by the scorer, if any. When provided, it is validated against ``TARGET_REQUIREMENTS``. """ - self._validator = validator if chat_target is not None: type(self).TARGET_REQUIREMENTS.validate(target=chat_target) @@ -203,173 +182,46 @@ def _create_identifier( async def score_async( self, - message: Message | None = None, *, - scorable: Scorable | None = None, + scorable: Scorable, expectation: ScoringExpectation | None = None, - objective: str | None = None, - role_filter: ChatMessageRole | None = None, - skip_on_error_result: bool = False, - infer_objective_from_request: bool = False, ) -> list[Score]: """ Score a scorable against an expectation, persist the results, and return them. - A scorer takes two inputs: the ``scorable`` says what to look at, and the - ``expectation`` says what to look for. Keeping them separate is what lets an attack - forward a question it does not understand to a scorer that does. - - Args: - message (Message | None): Deprecated. Pass - ``scorable=MessageScorable.from_message(...)`` instead. - scorable (Scorable | None): What to look at. + scorable (Scorable): What to look at. expectation (ScoringExpectation | None): What to look for. Defaults to None. - objective (str | None): Deprecated. Pass - ``expectation=ScoringExpectation(objective=...)`` instead. - role_filter (ChatMessageRole | None): Deprecated. Set ``role_filter`` on the - message scorable instead. - skip_on_error_result (bool): Deprecated. Set ``skip_on_error_result`` on the - message scorable instead. - infer_objective_from_request (bool): Deprecated. Resolve the objective at the - call site and pass it on the expectation instead. Returns: list[Score]: A list of Score objects representing the results. Raises: - ValueError: If the scorable inputs are missing, duplicated, or combined with - parameters that do not apply to them. TypeError: If this scorer does not support this kind of scorable. """ - resolved_scorable, resolved_expectation = self._consolidate_legacy_inputs( - message=message, - scorable=scorable, - expectation=expectation, - objective=objective, - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - infer_objective_from_request=infer_objective_from_request, - ) + scores = await self._score_scorable_async(scorable=scorable, expectation=expectation) + return self._validate_and_persist_scores(scores=scores) - scores = await self._score_scorable_async( - scorable=resolved_scorable, - expectation=resolved_expectation, - infer_objective_from_request=infer_objective_from_request, - ) + def _validate_and_persist_scores(self, *, scores: list[Score]) -> list[Score]: + """ + Validate and persist non-empty scorer output. + Returns: + list[Score]: The original scores. + """ if not scores: return [] self.validate_return_scores(scores=scores) self._memory.add_scores_to_memory(scores=scores) - return scores - def _consolidate_legacy_inputs( - self, - *, - message: Message | None, - scorable: Scorable | None, - expectation: ScoringExpectation | None, - objective: str | None, - role_filter: ChatMessageRole | None, - skip_on_error_result: bool, - infer_objective_from_request: bool, - ) -> tuple[Scorable, ScoringExpectation | None]: - """ - Map the deprecated message-shaped parameters onto a scorable and an expectation. - - Returns: - tuple[Scorable, ScoringExpectation | None]: The resolved scoring inputs. - - Raises: - ValueError: If the inputs are missing, duplicated, or combined with parameters - that do not apply to them. - """ - if message is not None and scorable is not None: - raise ValueError("Pass either 'message' or 'scorable', not both.") - if message is None and scorable is None: - raise ValueError("Either 'message' or 'scorable' must be provided.") - if objective is not None and expectation is not None: - raise ValueError("Pass either 'objective' or 'expectation', not both.") - if scorable is not None and (role_filter is not None or skip_on_error_result): - raise ValueError( - "'role_filter' and 'skip_on_error_result' are fields on the message scorable. " - "Set them on the scorable instead of passing them to score_async." - ) - - uses_legacy_parameters = ( - message is not None - or objective is not None - or role_filter is not None - or skip_on_error_result - or infer_objective_from_request - ) - if uses_legacy_parameters: - print_deprecation_message( - old_item="Scorer.score_async(message=..., objective=..., role_filter=..., " - "skip_on_error_result=..., infer_objective_from_request=...)", - new_item="Scorer.score_async(scorable=..., expectation=...)", - removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, - ) - - if scorable is None: - scorable = self._scorable_for_message( - cast("Message", message), - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) - - if objective is not None: - expectation = ScoringExpectation(objective=objective) - - return scorable, expectation - - @staticmethod - def _scorable_for_message( - message: Message, - *, - role_filter: ChatMessageRole | None = None, - skip_on_error_result: bool = False, - ) -> Scorable: - """ - Describe an in-hand message as the scorable that names it. - - The caller asked about this message, not about everything stored alongside it, so this - maps to the exact message rather than widening to its conversation. A message built by - hand has nothing in memory to name, so a single unpersisted piece becomes loose content. - - This bridges the two places that still hold a ``Message`` rather than the scorable it - came from: the deprecated ``message`` parameter, and the true/false wrappers that - forward to an inner scorer. Both go away in phase 2, and this goes with them. - - Args: - message (Message): The message to describe. - role_filter (ChatMessageRole | None): Role the piece must have to be scored. - skip_on_error_result (bool): Whether to skip the message when it holds an error. - - Returns: - Scorable: A ``ContentScorable`` for a single unpersisted piece, else a - ``MessageScorable`` naming the pieces. - """ - pieces = message.message_pieces - if len(pieces) == 1 and pieces[0].not_in_memory: - return ContentScorable.from_message(message) - - return MessageScorable.from_message( - message, - role_filter=role_filter, - skip_on_error_result=skip_on_error_result, - ) - @abstractmethod async def _score_scorable_async( self, *, scorable: Scorable, expectation: ScoringExpectation | None, - infer_objective_from_request: bool = False, ) -> list[Score]: """ Score a scorable this scorer supports. @@ -383,167 +235,12 @@ async def _score_scorable_async( Args: scorable (Scorable): What to look at. expectation (ScoringExpectation | None): What to look for. - infer_objective_from_request (bool): Deprecated; resolve the objective from the - stored conversation when the expectation carries none. Raises: TypeError: If the scorer does not support this kind of scorable. """ raise NotImplementedError - @staticmethod - def _create_scoring_text_piece( - *, - piece: MessagePiece, - content: str, - response_error: PromptResponseError, - ) -> MessagePiece: - """ - Create a text-typed scoring view that retains the persisted piece identity. - - Returns: - A text piece for scorer consumption. - """ - return MessagePiece( - id=piece.id, - role=piece.api_role, - original_value=piece.original_value, - converted_value=content, - original_value_data_type=piece.original_value_data_type, - converted_value_data_type="text", - conversation_id=piece.conversation_id, - sequence=piece.sequence, - prompt_metadata=dict(piece.prompt_metadata), - converter_identifiers=list(piece.converter_identifiers), # type: ignore[arg-type] - response_error=response_error, - timestamp=piece.timestamp, - original_prompt_id=piece.original_prompt_id, - not_in_memory=piece.not_in_memory, - ) - - @classmethod - def _create_text_piece_from_blocked(cls, piece: MessagePiece) -> MessagePiece | None: - """ - Create a text-typed copy of a blocked MessagePiece using its partial content. - - The substitute preserves the original piece's id (so scores link back correctly), - sets converted_value to the partial content with converted_value_data_type="text", - and sets response_error="none" so scorer short-circuits (e.g., refusal scorer's - blocked check) do not fire. - - Args: - piece: A blocked MessagePiece with prompt_metadata["partial_content"]. - - Returns: - MessagePiece with text content, or None if partial content is empty. - """ - partial_content = str(piece.prompt_metadata.get("partial_content", "")) - if not partial_content: - return None - - return cls._create_scoring_text_piece( - piece=piece, - content=partial_content, - response_error="none", - ) - - @classmethod - def _create_text_piece_from_structured_refusal(cls, piece: MessagePiece) -> MessagePiece | None: - """ - Create a blocked text scoring view for an SDK-provided structured refusal. - - Returns: - A text scoring view, or ``None`` when the piece is not a structured refusal. - """ - refusal = piece.structured_refusal - if not refusal: - return None - return cls._create_scoring_text_piece( - piece=piece, - content=refusal, - response_error="blocked", - ) - - def _apply_structured_refusal_substitution(self, message: Message) -> Message: - """ - Expose structured refusal explanations as text while preserving blocked semantics. - - Returns: - A scoring message with structured refusals substituted, or the original message. - """ - substituted = False - new_pieces: list[MessagePiece] = [] - for piece in message.message_pieces: - substitute = self._create_text_piece_from_structured_refusal(piece) - if substitute: - new_pieces.append(substitute) - substituted = True - continue - new_pieces.append(piece) - - return Message(message_pieces=new_pieces) if substituted else message - - def _apply_blocked_content_substitution(self, message: Message) -> Message: - """ - Create a copy of the message where blocked pieces with partial content are substituted. - - Each blocked piece that has prompt_metadata["partial_content"] is replaced with a - text-typed copy (response_error="none", converted_value=partial_content). Non-blocked - pieces and blocked pieces without partial content are kept as-is. - - Args: - message: The original message potentially containing blocked pieces. - - Returns: - A new Message with substituted pieces, or the original if no substitution was needed. - """ - substituted = False - new_pieces: list[MessagePiece] = [] - for piece in message.message_pieces: - if piece.is_blocked() and "partial_content" in piece.prompt_metadata: - substitute = self._create_text_piece_from_blocked(piece) - if substitute: - new_pieces.append(substitute) - substituted = True - continue - new_pieces.append(piece) - - if not substituted: - return message - - return Message(message_pieces=new_pieces) - - @abstractmethod - def _build_fallback_score( - self, *, message: Message, objective: str | None, scorer_response_blocked: bool = False - ) -> list[Score]: - """ - Return neutral fallback ``Score`` objects when ``_score_async`` produced no scores. - - Called from ``score_async`` after ``_score_async`` returns an empty list and the - message still has pieces (e.g. the response was blocked, had an error, or no piece - matched the validator). Every ``Scorer`` subclass MUST implement this so that a - consistent "attack did not succeed" value is always returned and downstream - consumers do not need to special-case error handling. - - Most scorers return a single-element list (e.g. ``FloatScaleScorer`` returns - ``[Score(0.0)]`` and ``TrueFalseScorer`` returns ``[Score(False)]``). Scorers - whose normal output shape is multiple scores per message (e.g. one per category) - should return one fallback score per logical output slot so downstream consumers - iterating by shape continue to work on blocked / error input. - - Args: - message (Message): The (possibly substituted) message that was scored. - objective (str | None): The objective associated with this scoring call. - scorer_response_blocked (bool): When True, the fallback was triggered because the - scorer's *own* LLM response was blocked by content filtering (not the - target-under-test). Subclasses should reflect this in the rationale. - - Returns: - list[Score]: One or more fallback scores. Must not be empty. - """ - ... - @abstractmethod def validate_return_scores(self, scores: list[Score]) -> None: """ @@ -689,20 +386,38 @@ async def score_prompts_batch_async( Raises: ValueError: If objectives is not None and the number of objectives doesn't match the number of messages. + TypeError: If this is not a message scorer. """ if objectives is None: - objectives = [""] * len(messages) + resolved_objectives = [""] * len(messages) elif len(objectives) != len(messages): raise ValueError("The number of objectives must match the number of messages.") + else: + resolved_objectives = list(objectives) if len(messages) == 0: return [] - scorables = [ - MessageScorable.from_message(message, role_filter=role_filter, skip_on_error_result=skip_on_error_result) - for message in messages - ] - expectations = [ScoringExpectation(objective=objective) for objective in objectives] + from pyrit.score.message_scorer import ( + MessageScorer, + MessageScoringOptions, + extract_objective_from_previous_turn, + ) + + if not isinstance(self, MessageScorer): + raise TypeError("score_prompts_batch_async requires a MessageScorer.") + if infer_objective_from_request: + resolved_objectives = [ + objective or extract_objective_from_previous_turn(message=message, memory=self._memory) + for message, objective in zip(messages, resolved_objectives, strict=True) + ] + + scorables = [MessageScorable.from_message(message) for message in messages] + expectations = [ScoringExpectation(objective=objective) for objective in resolved_objectives] + message_options = MessageScoringOptions( + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) # Some scorers do not have an associated prompt target; batch helper validates RPM only when present prompt_target = getattr(self, "_prompt_target", None) @@ -712,7 +427,7 @@ async def score_prompts_batch_async( prompt_target=cast("PromptTarget", prompt_target), batch_size=batch_size, items_to_batch=[scorables, expectations], - infer_objective_from_request=infer_objective_from_request, + message_options=message_options, ) # results is a list[list[Score]] and needs to be flattened @@ -791,6 +506,37 @@ def _extract_objective_from_response(self, response: Message) -> str: ) return extract_objective_from_previous_turn(message=response, memory=self._memory) + @staticmethod + async def _score_response_with_scorer_async( + *, + scorer: Scorer, + response: Message, + expectation: ScoringExpectation, + role_filter: ChatMessageRole, + skip_on_error_result: bool, + ) -> list[Score]: + """ + Apply response-scoring policy without storing policy on the scorable. + + Returns: + list[Score]: Scores from the message scorer. + + Raises: + TypeError: If the scorer does not use the message-scoring contract. + """ + from pyrit.score.message_scorer import MessageScorer, MessageScoringOptions + + if not isinstance(scorer, MessageScorer): + raise TypeError("Response scoring helpers require MessageScorer instances.") + return await scorer.score_async( + scorable=MessageScorable.from_message(response), + expectation=expectation, + message_options=MessageScoringOptions( + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ), + ) + @staticmethod async def score_response_async( *, @@ -848,21 +594,23 @@ async def score_response_async( objective=objective, skip_on_error_result=skip_on_error_result, ) - obj_task = objective_scorer.score_async( - scorable=MessageScorable.from_message( - response, role_filter=role_filter, skip_on_error_result=skip_on_error_result - ), + obj_task = Scorer._score_response_with_scorer_async( + scorer=objective_scorer, + response=response, expectation=ScoringExpectation(objective=objective), + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, ) aux_scores, obj_scores = await asyncio.gather(aux_task, obj_task) result["auxiliary_scores"] = aux_scores result["objective_scores"] = obj_scores else: - obj_scores = await objective_scorer.score_async( - scorable=MessageScorable.from_message( - response, role_filter=role_filter, skip_on_error_result=skip_on_error_result - ), + obj_scores = await Scorer._score_response_with_scorer_async( + scorer=objective_scorer, + response=response, expectation=ScoringExpectation(objective=objective), + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, ) result["objective_scores"] = obj_scores return result @@ -896,12 +644,17 @@ async def score_response_multiple_scorers_async( if not scorers: return [] - # Create all scoring tasks, note TEMPORARY fix to prevent multi-piece responses from breaking scoring logic - scorable = MessageScorable.from_message( - response, role_filter=role_filter, skip_on_error_result=skip_on_error_result - ) expectation = ScoringExpectation(objective=objective) - tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in scorers] + tasks = [ + Scorer._score_response_with_scorer_async( + scorer=scorer, + response=response, + expectation=expectation, + role_filter=role_filter, + skip_on_error_result=skip_on_error_result, + ) + for scorer in scorers + ] if not tasks: return [] diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index 04c2332b8b..74e5bc95a6 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -7,7 +7,16 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + ContentScorable, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, +) from pyrit.score.float_scale.float_scale_score_aggregator import FloatScaleAggregatorFunc, FloatScaleScoreAggregator from pyrit.score.float_scale.float_scale_scorer import FloatScaleScorer from pyrit.score.score_utils import ORIGINAL_FLOAT_VALUE_KEY @@ -99,7 +108,11 @@ async def _score_async( list[Score]: A list containing a single true/false Score object based on the threshold comparison. """ scores = await self._scorer.score_async( - scorable=self._scorable_for_message(message, role_filter=role_filter), + scorable=( + ContentScorable.from_message(message) + if len(message.message_pieces) == 1 and message.get_piece().not_in_memory + else MessageScorable.from_message(message) + ), expectation=ScoringExpectation(objective=objective), ) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index 88538467cc..6236abd748 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -7,7 +7,16 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + ContentScorable, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, +) from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import TrueFalseAggregatorFunc from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -97,7 +106,11 @@ async def _score_async( ValueError: If any constituent scorer does not return exactly one score. ValueError: If no scores are generated from the request response pieces. """ - scorable = self._scorable_for_message(message, role_filter=role_filter) + scorable = ( + ContentScorable.from_message(message) + if len(message.message_pieces) == 1 and message.get_piece().not_in_memory + else MessageScorable.from_message(message) + ) expectation = ScoringExpectation(objective=objective) tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in self._scorers] diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index 7184642fa9..0fdec6652c 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -7,7 +7,16 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + ContentScorable, + Message, + MessagePiece, + MessageScorable, + Score, + ScoringExpectation, +) from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_scorer import TrueFalseScorer @@ -74,7 +83,11 @@ async def _score_async( list[Score]: A list containing a single Score object with the inverted true/false value. """ scores = await self._scorer.score_async( - scorable=self._scorable_for_message(message, role_filter=role_filter), + scorable=( + ContentScorable.from_message(message) + if len(message.message_pieces) == 1 and message.get_piece().not_in_memory + else MessageScorable.from_message(message) + ), expectation=ScoringExpectation(objective=objective), ) inv_score = scores[0] diff --git a/pyrit/score/true_false/true_false_scorer.py b/pyrit/score/true_false/true_false_scorer.py index 36524dcbc9..298a37f95a 100644 --- a/pyrit/score/true_false/true_false_scorer.py +++ b/pyrit/score/true_false/true_false_scorer.py @@ -11,6 +11,7 @@ if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget + from pyrit.score.message_scorable_resolver import MessageScorableResolver from pyrit.score.scorer_evaluation.scorer_evaluator import ScorerEvalDatasetFiles from pyrit.score.scorer_evaluation.scorer_metrics import ObjectiveScorerMetrics from pyrit.score.scorer_prompt_validator import ScorerPromptValidator @@ -47,6 +48,7 @@ def __init__( validator: ScorerPromptValidator, score_aggregator: TrueFalseAggregatorFunc = TrueFalseScoreAggregator.OR, chat_target: PromptTarget | None = None, + message_resolver: MessageScorableResolver | None = None, ) -> None: """ Initialize the TrueFalseScorer. @@ -57,6 +59,7 @@ def __init__( Defaults to TrueFalseScoreAggregator.OR. chat_target (PromptTarget | None): Optional chat target used by the scorer, forwarded to the base class for validation against ``TARGET_REQUIREMENTS``. + message_resolver (MessageScorableResolver | None): Message evidence resolver. """ self._score_aggregator = score_aggregator @@ -69,7 +72,11 @@ def __init__( result_file="objective/objective_achieved_metrics.jsonl", ) - super().__init__(validator=validator, chat_target=chat_target) + super().__init__( + validator=validator, + chat_target=chat_target, + message_resolver=message_resolver, + ) def validate_return_scores(self, scores: list[Score]) -> None: """ diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo.py b/tests/unit/executor/attack/multi_turn/test_crescendo.py index d0f3c865f0..c599b13e18 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo.py @@ -1201,10 +1201,10 @@ async def test_check_refusal_does_not_skip_on_error_result( await attack._check_refusal_async(context=basic_context, objective="test task") - # Verify the scorable does not skip error results + # Verify message policy does not skip error results mock_refusal_scorer.score_async.assert_called_once() - scorable = mock_refusal_scorer.score_async.call_args.kwargs["scorable"] - assert scorable.skip_on_error_result is False, ( + message_options = mock_refusal_scorer.score_async.call_args.kwargs["message_options"] + assert message_options.skip_on_error_result is False, ( "Refusal scorer must be called with skip_on_error_result=False " "to ensure error responses are scored (treated as refusals) rather than skipped" ) diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py index ba2e604990..648b4c1c97 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo_resilience.py @@ -23,6 +23,7 @@ from pyrit.models import AttackOutcome, ComponentIdentifier, ConversationType, Message, MessagePiece, Score from pyrit.prompt_normalizer import PromptNormalizer from pyrit.score import Scorer, TrueFalseScorer +from pyrit.score.message_scorable_resolver import MessageScorableResolver _OBJECTIVE = "Recover the hidden phrase through gradual rapport." @@ -283,7 +284,7 @@ async def score_objective(**_kwargs): # A scorable names piece ids rather than carrying the message, so read them back. memory = CentralMemory.get_memory_instance() refusal_inputs = [ - call.kwargs["scorable"].resolve_message(memory=memory).get_value() + MessageScorableResolver().resolve(scorable=call.kwargs["scorable"], memory=memory).get_value() for call in refusal_scorer.score_async.await_args_list ] assert refusal_inputs == [ diff --git a/tests/unit/models/test_scorable.py b/tests/unit/models/test_scorable.py new file mode 100644 index 0000000000..006a7f5ec2 --- /dev/null +++ b/tests/unit/models/test_scorable.py @@ -0,0 +1,86 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import dataclasses +import uuid + +import pytest + +from pyrit.models import ContentScorable, Message, MessagePiece, MessageScorable, Scorable + + +def _message(value: str = "response") -> Message: + return MessagePiece( + role="assistant", + original_value=value, + conversation_id=str(uuid.uuid4()), + ).to_message() + + +@pytest.mark.parametrize( + "scorable, field_name", + [ + (MessageScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), + (ContentScorable(value="hello"), "value"), + ], +) +def test_scorable_is_frozen(scorable: Scorable, field_name: str): + with pytest.raises(dataclasses.FrozenInstanceError): + setattr(scorable, field_name, "changed") + + +@pytest.mark.parametrize( + "scorable", + [ + MessageScorable(message_piece_ids=(uuid.uuid4(),)), + ContentScorable(value="hello"), + ], +) +def test_every_scorable_is_a_scorable(scorable: Scorable): + assert isinstance(scorable, Scorable) + + +def test_scorables_are_inert(): + assert not hasattr(MessageScorable(message_piece_ids=(uuid.uuid4(),)), "resolve_message") + assert not hasattr(ContentScorable(value="hello"), "to_ephemeral_message") + + +def test_scorables_are_keyword_only(): + with pytest.raises(TypeError): + ContentScorable("hello") # type: ignore[misc] + + +def test_message_scorable_defaults(): + piece_id = uuid.uuid4() + + scorable = MessageScorable(message_piece_ids=(piece_id,)) + + assert scorable.message_piece_ids == (piece_id,) + + +def test_message_scorable_from_message_names_pieces(): + message = _message() + + scorable = MessageScorable.from_message(message) + + assert scorable.message_piece_ids == (message.get_piece().id,) + assert not hasattr(scorable, "message") + + +def test_content_scorable_defaults_to_text(): + assert ContentScorable(value="hello").data_type == "text" + + +def test_content_scorable_from_message_uses_converted_view(): + message = MessagePiece( + role="user", + original_value="original", + converted_value="converted", + original_value_data_type="text", + converted_value_data_type="text", + ).to_message() + + scorable = ContentScorable.from_message(message) + + assert scorable.value == "converted" + assert scorable.data_type == "text" diff --git a/tests/unit/score/test_message_scorable_resolver.py b/tests/unit/score/test_message_scorable_resolver.py new file mode 100644 index 0000000000..4abba8330a --- /dev/null +++ b/tests/unit/score/test_message_scorable_resolver.py @@ -0,0 +1,95 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import uuid +from unittest.mock import MagicMock + +import pytest + +from pyrit.memory import MemoryInterface +from pyrit.models import ContentScorable, Message, MessagePiece, MessageScorable +from pyrit.score.message_scorable_resolver import MessageScorableResolver + + +def _stored_message(value: str = "stored response") -> Message: + return MessagePiece( + role="assistant", + original_value=value, + conversation_id=str(uuid.uuid4()), + ).to_message() + + +def test_resolver_reads_message_reference_from_memory(sqlite_instance: MemoryInterface): + stored = _stored_message() + sqlite_instance.add_message_to_memory(request=stored) + + resolved = MessageScorableResolver().resolve( + scorable=MessageScorable.from_message(stored), + memory=sqlite_instance, + ) + + assert resolved.get_value() == "stored response" + + +def test_resolver_reports_missing_piece_ids(sqlite_instance: MemoryInterface): + stored = _stored_message() + sqlite_instance.add_message_to_memory(request=stored) + missing_id = uuid.uuid4() + + with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): + MessageScorableResolver().resolve( + scorable=MessageScorable(message_piece_ids=(stored.get_piece().id, missing_id)), + memory=sqlite_instance, + ) + + +def test_resolver_rejects_pieces_from_multiple_messages(sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + first = MessagePiece( + role="user", + original_value="ask", + conversation_id=conversation_id, + sequence=0, + ).to_message() + second = MessagePiece( + role="assistant", + original_value="answer", + conversation_id=conversation_id, + sequence=1, + ).to_message() + sqlite_instance.add_message_to_memory(request=first) + sqlite_instance.add_message_to_memory(request=second) + + with pytest.raises(ValueError, match="exactly one message"): + MessageScorableResolver().resolve( + scorable=MessageScorable( + message_piece_ids=(first.get_piece().id, second.get_piece().id), + ), + memory=sqlite_instance, + ) + + +def test_resolver_preserves_reference_order(sqlite_instance: MemoryInterface): + conversation_id = str(uuid.uuid4()) + first = MessagePiece(role="assistant", original_value="one", conversation_id=conversation_id, sequence=0) + second = MessagePiece(role="assistant", original_value="two", conversation_id=conversation_id, sequence=0) + sqlite_instance.add_message_to_memory(request=Message(message_pieces=[first, second])) + + resolved = MessageScorableResolver().resolve( + scorable=MessageScorable(message_piece_ids=(second.id, first.id)), + memory=sqlite_instance, + ) + + assert [piece.original_value for piece in resolved.message_pieces] == ["two", "one"] + + +def test_resolver_adapts_content_to_ephemeral_message(): + resolved = MessageScorableResolver().resolve( + scorable=ContentScorable(value="loose text"), + memory=MagicMock(spec=MemoryInterface), + ) + + piece = resolved.get_piece() + assert piece.converted_value == "loose text" + assert piece.role == "user" + assert piece.not_in_memory is True diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py index 7f263e7071..de6f4a40f7 100644 --- a/tests/unit/score/test_message_scorer.py +++ b/tests/unit/score/test_message_scorer.py @@ -2,7 +2,9 @@ # Licensed under the MIT license. import dataclasses +import inspect import uuid +from unittest.mock import MagicMock import pytest @@ -17,7 +19,8 @@ ScorerPromptValidator, TrueFalseScorer, ) -from pyrit.score.message_scorer import extract_objective_from_previous_turn +from pyrit.score.message_scorable_resolver import MessageScorableResolver +from pyrit.score.message_scorer import MessageScoringOptions, extract_objective_from_previous_turn @dataclasses.dataclass(frozen=True) @@ -38,8 +41,8 @@ def is_message_piece_supported(self, message_piece): class RecordingScorer(TrueFalseScorer): """A message scorer that remembers what it was asked to score.""" - def __init__(self): - super().__init__(validator=PermissiveValidator()) + def __init__(self, *, message_resolver: MessageScorableResolver | None = None): + super().__init__(validator=PermissiveValidator(), message_resolver=message_resolver) self.scored_messages: list[Message] = [] self.scored_objectives: list[str | None] = [] @@ -165,6 +168,16 @@ async def test_content_scorable_is_never_persisted(self): # Memory cannot link a score to a piece it never stored. assert scores[0].message_piece_id is None + async def test_message_scorer_uses_injected_resolver(self): + message = _assistant_message() + resolver = MagicMock(spec=MessageScorableResolver) + resolver.resolve.return_value = message + scorer = RecordingScorer(message_resolver=resolver) + + await scorer.score_async(scorable=ContentScorable(value="ignored")) + + resolver.resolve.assert_called_once() + async def test_unsupported_scorable_raises_type_error(self): scorer = RecordingScorer() @@ -173,7 +186,7 @@ async def test_unsupported_scorable_raises_type_error(self): class TestScorerBaseIsScorableAgnostic: - """Scorer knows nothing about messages; the message hooks belong to MessageScorer.""" + """The base extension contract contains no message-processing requirements.""" def test_scorer_requires_a_scorable_implementation(self): # A scorer that implements only the message hooks cannot be instantiated. Without @@ -189,16 +202,29 @@ def test_message_scorer_satisfies_the_scorable_contract(self): assert "_score_scorable_async" not in MessageScorer.__abstractmethods__ assert "_score_piece_async" in MessageScorer.__abstractmethods__ + def test_message_dependencies_live_on_message_scorer(self): + assert "validator" not in inspect.signature(Scorer).parameters + for hook in [ + "_build_fallback_score", + "_apply_structured_refusal_substitution", + "_apply_blocked_content_substitution", + ]: + assert not hasattr(Scorer, hook) + assert hasattr(MessageScorer, hook) + @pytest.mark.usefixtures("patch_central_database") class TestScorableFilters: - """role_filter and skip_on_error_result are fields on the scorable, not call parameters.""" + """Message policy is separate from the scorable's evidence identity.""" async def test_role_filter_mismatch_skips_scoring(self): scorer = RecordingScorer() message = _assistant_message() - scores = await scorer.score_async(scorable=MessageScorable.from_message(message, role_filter="user")) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(message), + message_options=MessageScoringOptions(role_filter="user"), + ) assert scores == [] assert scorer.scored_messages == [] @@ -207,7 +233,10 @@ async def test_role_filter_match_scores(self): scorer = RecordingScorer() message = _assistant_message() - scores = await scorer.score_async(scorable=MessageScorable.from_message(message, role_filter="assistant")) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(message), + message_options=MessageScoringOptions(role_filter="assistant"), + ) assert len(scores) == 1 @@ -215,7 +244,10 @@ async def test_skip_on_error_result_skips_error_message(self): scorer = RecordingScorer() message = _error_message() - scores = await scorer.score_async(scorable=MessageScorable.from_message(message, skip_on_error_result=True)) + scores = await scorer.score_async( + scorable=MessageScorable.from_message(message), + message_options=MessageScoringOptions(skip_on_error_result=True), + ) assert scores == [] assert scorer.scored_messages == [] @@ -274,6 +306,20 @@ async def test_keyword_message_maps_to_message_scorable(self): assert scorer.scored_messages == [message] + async def test_ephemeral_message_maps_its_converted_view_to_content(self): + scorer = RecordingScorer() + message = MessagePiece( + role="user", + original_value="original", + converted_value="converted", + ).to_message() + message.set_response_not_in_memory() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async(message) + + assert scorer.scored_messages[0].get_value() == "converted" + async def test_message_does_not_widen_to_the_stored_conversation(self, sqlite_instance: MemoryInterface): """The shim scores the supplied message, never the whole conversation behind it.""" conversation_id = str(uuid.uuid4()) @@ -301,7 +347,7 @@ async def test_objective_maps_to_expectation(self): assert scorer.scored_objectives == ["legacy objective"] - async def test_legacy_role_filter_maps_onto_the_scorable(self): + async def test_legacy_role_filter_maps_to_message_options(self): scorer = RecordingScorer() with pytest.warns(DeprecationWarning, match="Scorer.score_async"): @@ -309,7 +355,7 @@ async def test_legacy_role_filter_maps_onto_the_scorable(self): assert scores == [] - async def test_legacy_skip_on_error_result_maps_onto_the_scorable(self): + async def test_legacy_skip_on_error_result_maps_to_message_options(self): scorer = RecordingScorer() message = _error_message() @@ -318,6 +364,22 @@ async def test_legacy_skip_on_error_result_maps_onto_the_scorable(self): assert scores == [] + @pytest.mark.parametrize( + "kwargs", + [ + {"skip_on_error_result": False}, + {"infer_objective_from_request": False}, + ], + ) + async def test_explicit_false_legacy_boolean_emits_warning(self, kwargs): + scorer = RecordingScorer() + + with pytest.warns(DeprecationWarning, match="Scorer.score_async"): + await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + **kwargs, + ) + async def test_infer_objective_from_request_reads_the_previous_turn(self, sqlite_instance: MemoryInterface): conversation_id = str(uuid.uuid4()) sqlite_instance.add_message_to_memory( @@ -375,11 +437,15 @@ async def test_objective_and_expectation_together_raises(self): ) @pytest.mark.parametrize("kwargs", [{"role_filter": "assistant"}, {"skip_on_error_result": True}]) - async def test_message_flags_with_a_scorable_raises(self, kwargs): + async def test_message_options_and_legacy_policy_raise(self, kwargs): scorer = RecordingScorer() - with pytest.raises(ValueError, match="fields on the message scorable"): - await scorer.score_async(scorable=MessageScorable.from_message(_assistant_message()), **kwargs) + with pytest.raises(ValueError, match="either 'message_options' or legacy"): + await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + message_options=MessageScoringOptions(), + **kwargs, + ) @pytest.mark.usefixtures("patch_central_database") diff --git a/tests/unit/score/test_scorable.py b/tests/unit/score/test_scorable.py deleted file mode 100644 index 369d5774ba..0000000000 --- a/tests/unit/score/test_scorable.py +++ /dev/null @@ -1,175 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -import dataclasses -import uuid - -import pytest - -from pyrit.memory import MemoryInterface -from pyrit.models import Message, MessagePiece -from pyrit.score import ContentScorable, MessageScorable, Scorable - - -def _stored_message(value: str = "stored response"): - """Return a message memory will accept, so it needs a conversation id.""" - return MessagePiece( - role="assistant", - original_value=value, - conversation_id=str(uuid.uuid4()), - ).to_message() - - -@pytest.mark.parametrize( - "scorable, field_name", - [ - (MessageScorable(message_piece_ids=(uuid.uuid4(),)), "message_piece_ids"), - (ContentScorable(value="hello"), "value"), - ], -) -def test_scorable_is_frozen(scorable, field_name): - with pytest.raises(dataclasses.FrozenInstanceError): - setattr(scorable, field_name, "changed") - - -@pytest.mark.parametrize( - "scorable", - [ - MessageScorable(message_piece_ids=(uuid.uuid4(),)), - ContentScorable(value="hello"), - ], -) -def test_every_scorable_is_a_scorable(scorable): - assert isinstance(scorable, Scorable) - - -def test_content_scorable_names_content_rather_than_a_message(): - """Loose content has no message behind it, so it must not answer the message contract.""" - scorable = ContentScorable(value="hello") - - assert isinstance(scorable, Scorable) - assert not isinstance(scorable, MessageScorable) - assert not hasattr(scorable, "resolve_message") - - -@pytest.mark.parametrize("field_name", ["role_filter", "skip_on_error_result"]) -def test_content_scorable_rejects_message_filters(field_name): - # Loose content has no role and no error state. Accepting a role filter here used to - # score nothing at all, silently, because the adapted message is always role="user". - with pytest.raises(TypeError): - ContentScorable(value="hello", **{field_name: "assistant"}) - - -def test_scorables_are_keyword_only(): - with pytest.raises(TypeError): - ContentScorable("hello") # type: ignore[misc] - - -class TestMessageScorable: - def test_defaults(self): - piece_id = uuid.uuid4() - scorable = MessageScorable(message_piece_ids=(piece_id,)) - - assert scorable.message_piece_ids == (piece_id,) - assert scorable.role_filter is None - assert scorable.skip_on_error_result is False - - def test_from_message_names_the_pieces_rather_than_carrying_them(self): - """A scorable is a reference, so a message in hand becomes its piece ids.""" - message = _stored_message() - - scorable = MessageScorable.from_message(message, role_filter="assistant") - - assert scorable.message_piece_ids == (message.get_piece().id,) - assert scorable.role_filter == "assistant" - assert not hasattr(scorable, "message") - - def test_resolves_from_memory(self, sqlite_instance: MemoryInterface): - stored = _stored_message() - sqlite_instance.add_message_to_memory(request=stored) - scorable = MessageScorable.from_message(stored) - - assert scorable.resolve_message(memory=sqlite_instance).get_value() == "stored response" - - def test_raises_when_nothing_is_in_memory(self, sqlite_instance: MemoryInterface): - missing_id = uuid.uuid4() - scorable = MessageScorable(message_piece_ids=(missing_id,)) - - with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): - scorable.resolve_message(memory=sqlite_instance) - - def test_raises_naming_only_the_missing_ids(self, sqlite_instance: MemoryInterface): - """A partial resolution is a caller error, and the error must point at the bad ids.""" - stored = _stored_message() - sqlite_instance.add_message_to_memory(request=stored) - missing_id = uuid.uuid4() - scorable = MessageScorable(message_piece_ids=(stored.get_piece().id, missing_id)) - - with pytest.raises(ValueError, match=f"No message pieces found in memory for ids \\['{missing_id}'\\]"): - scorable.resolve_message(memory=sqlite_instance) - - def test_raises_when_pieces_span_more_than_one_message(self, sqlite_instance: MemoryInterface): - conversation_id = str(uuid.uuid4()) - first = MessagePiece( - role="user", original_value="ask", conversation_id=conversation_id, sequence=0 - ).to_message() - second = MessagePiece( - role="assistant", original_value="answer", conversation_id=conversation_id, sequence=1 - ).to_message() - sqlite_instance.add_message_to_memory(request=first) - sqlite_instance.add_message_to_memory(request=second) - scorable = MessageScorable(message_piece_ids=(first.get_piece().id, second.get_piece().id)) - - with pytest.raises(ValueError, match="exactly one message"): - scorable.resolve_message(memory=sqlite_instance) - - -class TestContentScorable: - def test_defaults_to_text(self): - assert ContentScorable(value="hello").data_type == "text" - - def test_adapts_to_an_unpersisted_message(self): - """Loose scoring must stay usable with no memory behind it.""" - message = ContentScorable(value="loose text").to_ephemeral_message() - - piece = message.get_piece() - assert piece.original_value == "loose text" - assert piece.role == "user" - assert piece.not_in_memory is True - - def test_adapts_non_text_data_types(self): - message = ContentScorable(value="path/to.png", data_type="image_path").to_ephemeral_message() - - assert message.get_piece().original_value_data_type == "image_path" - - -class TestFromMessage: - """A caller holding a persisted Message names its pieces rather than carrying it.""" - - def test_persisted_message_becomes_a_reference(self): - message = _stored_message() - - scorable = MessageScorable.from_message(message) - - assert isinstance(scorable, MessageScorable) - assert scorable.message_piece_ids == (message.get_piece().id,) - - def test_filters_travel_onto_the_reference(self): - scorable = MessageScorable.from_message(_stored_message(), role_filter="assistant", skip_on_error_result=True) - - assert isinstance(scorable, MessageScorable) - assert scorable.role_filter == "assistant" - assert scorable.skip_on_error_result is True - - -def test_resolution_preserves_the_order_the_scorable_names(sqlite_instance: MemoryInterface): - """Memory returns its own order, but a multi-piece message reads differently if shuffled.""" - conversation_id = str(uuid.uuid4()) - first = MessagePiece(role="assistant", original_value="one", conversation_id=conversation_id, sequence=0) - second = MessagePiece(role="assistant", original_value="two", conversation_id=conversation_id, sequence=0) - sqlite_instance.add_message_to_memory(request=Message(message_pieces=[first, second])) - - reversed_scorable = MessageScorable(message_piece_ids=(second.id, first.id)) - - resolved = reversed_scorable.resolve_message(memory=sqlite_instance) - assert [piece.original_value for piece in resolved.message_pieces] == ["two", "one"] diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index 2a85d84fe2..2ee17ccf26 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -23,7 +23,8 @@ TrueFalseScorer, ) from pyrit.score.llm_scoring import _run_llm_scoring_async -from pyrit.score.message_scorer import extract_objective_from_previous_turn +from pyrit.score.message_scorable_resolver import MessageScorableResolver +from pyrit.score.message_scorer import MessageScoringOptions, extract_objective_from_previous_turn @pytest.fixture @@ -458,6 +459,7 @@ async def test_scorer_score_responses_batch_async(patch_central_database): assert first_call_kwargs["scorable"] == MessageScorable.from_message(store_message(user_req)) assert first_call_kwargs["expectation"] == ScoringExpectation(objective="") + assert first_call_kwargs["message_options"] == MessageScoringOptions() assert fake_scores[0] in results assert len(fake_scores) == 2 @@ -568,16 +570,17 @@ async def test_score_response_async_parallel_execution(patch_central_database): assert score1_1 in result["auxiliary_scores"] assert score2_1 in result["auxiliary_scores"] - expected_scorable = MessageScorable.from_message( - store_message(response), role_filter="assistant", skip_on_error_result=True - ) + expected_scorable = MessageScorable.from_message(store_message(response)) + expected_options = MessageScoringOptions(role_filter="assistant", skip_on_error_result=True) scorer1.score_async.assert_any_call( scorable=expected_scorable, expectation=ScoringExpectation(objective="test task"), + message_options=expected_options, ) scorer2.score_async.assert_any_call( scorable=expected_scorable, expectation=ScoringExpectation(objective="test task"), + message_options=expected_options, ) @@ -599,8 +602,9 @@ async def test_score_async_no_matching_role(patch_central_database): response = Message(message_pieces=[MessagePiece(role="user", original_value="test", conversation_id="test-convo")]) scorer = MockScorer() result = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(response), role_filter="assistant"), + scorable=MessageScorable.from_message(store_message(response)), expectation=ScoringExpectation(objective="test task"), + message_options=MessageScoringOptions(role_filter="assistant"), ) assert result == [] @@ -694,7 +698,15 @@ async def test_score_response_success_async_parallel_scoring_per_piece(patch_cen def _first_value(scorable: MessageScorable) -> str: # A scorable names pieces rather than carrying them, so read it back from memory. - return scorable.resolve_message(memory=CentralMemory.get_memory_instance()).message_pieces[0].original_value + return ( + MessageScorableResolver() + .resolve( + scorable=scorable, + memory=CentralMemory.get_memory_instance(), + ) + .message_pieces[0] + .original_value + ) async def mock_score_async_1(*, scorable: MessageScorable, **kwargs) -> list[Score]: call_order.append(("scorer1", _first_value(scorable))) @@ -2023,7 +2035,7 @@ def _make_normal_piece(*, conversation_id: str = "test-convo") -> MessagePiece: class TestCreateTextPieceFromBlocked: def test_returns_text_piece_with_partial_content(self): piece = _make_blocked_piece(partial_content="Harmful partial text here") - substitute = Scorer._create_text_piece_from_blocked(piece) + substitute = MessageScorer._create_text_piece_from_blocked(piece) assert substitute is not None assert substitute.converted_value == "Harmful partial text here" @@ -2033,7 +2045,7 @@ def test_returns_text_piece_with_partial_content(self): def test_preserves_original_value(self): piece = _make_blocked_piece(partial_content="partial") - substitute = Scorer._create_text_piece_from_blocked(piece) + substitute = MessageScorer._create_text_piece_from_blocked(piece) assert substitute is not None assert substitute.original_value == piece.original_value @@ -2041,22 +2053,22 @@ def test_preserves_original_value(self): def test_returns_none_when_no_partial_content(self): piece = _make_blocked_piece() - assert Scorer._create_text_piece_from_blocked(piece) is None + assert MessageScorer._create_text_piece_from_blocked(piece) is None def test_returns_none_when_empty_partial_content(self): piece = _make_blocked_piece(partial_content="") - assert Scorer._create_text_piece_from_blocked(piece) is None + assert MessageScorer._create_text_piece_from_blocked(piece) is None def test_preserves_conversation_id(self): piece = _make_blocked_piece(partial_content="partial") - substitute = Scorer._create_text_piece_from_blocked(piece) + substitute = MessageScorer._create_text_piece_from_blocked(piece) assert substitute is not None assert substitute.conversation_id == piece.conversation_id def test_response_error_is_none_not_blocked(self): """Substitute must have response_error='none' so refusal short-circuits don't fire.""" piece = _make_blocked_piece(partial_content="partial text") - substitute = Scorer._create_text_piece_from_blocked(piece) + substitute = MessageScorer._create_text_piece_from_blocked(piece) assert substitute is not None assert substitute.response_error == "none" assert not substitute.is_blocked() @@ -2067,7 +2079,7 @@ class TestCreateTextPieceFromStructuredRefusal: def test_returns_blocked_text_piece_with_refusal_explanation(self): piece = _make_blocked_piece(structured_refusal="I cannot assist with that request.") - substitute = Scorer._create_text_piece_from_structured_refusal(piece) + substitute = MessageScorer._create_text_piece_from_structured_refusal(piece) assert substitute is not None assert substitute.converted_value == "I cannot assist with that request." @@ -2076,7 +2088,7 @@ def test_returns_blocked_text_piece_with_refusal_explanation(self): assert substitute.id == piece.id def test_returns_none_for_generic_blocked_response(self): - assert Scorer._create_text_piece_from_structured_refusal(_make_blocked_piece()) is None + assert MessageScorer._create_text_piece_from_structured_refusal(_make_blocked_piece()) is None # ── score_async with score_blocked_content tests ───────────────────────────── @@ -2182,7 +2194,8 @@ async def test_skip_on_error_true_without_flag_skips_blocked(self): msg = Message(message_pieces=[_make_blocked_piece(partial_content="harmful text")]) scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert scores == [] @@ -2192,7 +2205,8 @@ async def test_skip_on_error_true_with_flag_does_not_skip_when_partial_content(s scorer.score_blocked_content = True scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert len(scores) == 1 assert scores[0].score_value == "true" @@ -2203,7 +2217,8 @@ async def test_skip_on_error_true_with_flag_still_skips_when_no_partial_content( scorer.score_blocked_content = True scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert scores == [] @@ -2222,7 +2237,8 @@ async def test_skip_on_error_skips_error_type_without_response_error_flag(self): ) scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert scores == [] @@ -2242,7 +2258,8 @@ async def test_skip_on_error_scores_structured_refusal_as_text(self, validator: msg = Message(message_pieces=[piece]) scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert len(scores) == 1 @@ -2272,7 +2289,8 @@ async def test_skip_on_error_still_skips_mixed_structured_and_runtime_errors(sel ) scores = await scorer.score_async( - scorable=MessageScorable.from_message(store_message(msg), skip_on_error_result=True) + scorable=MessageScorable.from_message(store_message(msg)), + message_options=MessageScoringOptions(skip_on_error_result=True), ) assert scores == [] From 433ca00bd6a877ef93a3ef4dc8088a911d9a4271 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Wed, 19 Aug 2026 13:03:44 -0700 Subject: [PATCH 09/11] Add the conditions axis and MatchesObjective to the expectation --- pyrit/models/__init__.py | 4 +++ pyrit/models/score/__init__.py | 3 ++ pyrit/models/score/condition.py | 28 +++++++++++++++++ pyrit/models/score/expectation.py | 21 ++++++++++--- tests/unit/models/test_expectation.py | 44 +++++++++++++++++++++++++-- 5 files changed, 93 insertions(+), 7 deletions(-) create mode 100644 pyrit/models/score/condition.py diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index ca03c5c7e1..667a2ca033 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -79,7 +79,9 @@ from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT from pyrit.models.retry_event import RetryEvent from pyrit.models.score import ( + Condition, ContentScorable, + MatchesObjective, MessageScorable, Scorable, Score, @@ -140,6 +142,7 @@ "ComponentIdentifier", "ComponentType", "compute_eval_hash", + "Condition", "config_hash", "ConverterIdentifier", "Conversation", @@ -176,6 +179,7 @@ "JSON_SCHEMA_METADATA_KEY", "SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY", "JsonSchemaDefinition", + "MatchesObjective", "MEDIA_PATH_DATA_TYPES", "Message", "MessagePiece", diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py index 28151b0d4b..b870e813b9 100644 --- a/pyrit/models/score/__init__.py +++ b/pyrit/models/score/__init__.py @@ -9,13 +9,16 @@ are inert canonical data; scoring-layer resolvers acquire the evidence they name. """ +from pyrit.models.score.condition import Condition, MatchesObjective from pyrit.models.score.expectation import ScoringExpectation from pyrit.models.score.scorable import ContentScorable, MessageScorable, Scorable from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore __all__ = [ "ComponentIdentifierField", + "Condition", "ContentScorable", + "MatchesObjective", "MessageScorable", "Scorable", "Score", diff --git a/pyrit/models/score/condition.py b/pyrit/models/score/condition.py new file mode 100644 index 0000000000..25ddaeedad --- /dev/null +++ b/pyrit/models/score/condition.py @@ -0,0 +1,28 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from abc import ABC +from dataclasses import dataclass + + +class Condition(ABC): # noqa: B024 root type; each scoring domain declares its own criterion + """ + What counts as satisfied. + + A condition is a neutral predicate about evidence: it says what to detect, never + whether detecting it is good or bad. Polarity belongs to a scorer that wraps another, + such as ``TrueFalseInverterScorer``. Each scoring domain adds its own subclass. + """ + + +@dataclass(frozen=True, kw_only=True) +class MatchesObjective(Condition): + """ + The evidence satisfies the expectation's own objective, as a judge reads it. + + This carries no text of its own. The objective lives on the ``ScoringExpectation``, + so a scorer matching this condition reads it from there and the two can never + disagree. + """ diff --git a/pyrit/models/score/expectation.py b/pyrit/models/score/expectation.py index 924f52d487..b795a17855 100644 --- a/pyrit/models/score/expectation.py +++ b/pyrit/models/score/expectation.py @@ -3,17 +3,28 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field +from pyrit.models.score.condition import Condition # noqa: TC001 (runtime-required by dataclass field annotations) -@dataclass(frozen=True) + +@dataclass(frozen=True, kw_only=True) class ScoringExpectation: """ What a scorer scores against. - An expectation is a single parameter that attacks forward without inspecting it, - so a question authored in a technique configuration or a seed can reach a scorer - through an attack that knows nothing about it. + An expectation is a single parameter, so a question authored in a technique + configuration or a seed can reach a scorer through an attack that knows nothing + about it. It has two independent axes. + + ``objective`` carries the intent: prose describing what the run is trying to do. + Components read it for framing — an adversarial target renders it into a system + prompt, a report prints it — and none of them match it. + + ``conditions`` carry the criteria: typed objects routed by type to the scorers that + match them. Attacks forward them without inspecting them, and a scorer matches at + most one of them. """ objective: str | None = None + conditions: tuple[Condition, ...] = field(default_factory=tuple) diff --git a/tests/unit/models/test_expectation.py b/tests/unit/models/test_expectation.py index 8635d0622a..2fc5a2e2df 100644 --- a/tests/unit/models/test_expectation.py +++ b/tests/unit/models/test_expectation.py @@ -5,11 +5,14 @@ import pytest -from pyrit.models import ScoringExpectation +from pyrit.models import Condition, MatchesObjective, ScoringExpectation def test_expectation_defaults(): - assert ScoringExpectation().objective is None + expectation = ScoringExpectation() + + assert expectation.objective is None + assert expectation.conditions == () def test_expectation_is_frozen(): @@ -21,3 +24,40 @@ def test_expectation_is_frozen(): def test_expectations_with_equal_values_compare_equal(): assert ScoringExpectation(objective="a") == ScoringExpectation(objective="a") + + +def test_expectation_carries_conditions_beside_the_objective(): + expectation = ScoringExpectation(objective="exfiltrate", conditions=(MatchesObjective(),)) + + assert expectation.objective == "exfiltrate" + assert expectation.conditions == (MatchesObjective(),) + + +def test_expectation_carries_conditions_without_an_objective(): + expectation = ScoringExpectation(conditions=(MatchesObjective(),)) + + assert expectation.objective is None + assert expectation.conditions == (MatchesObjective(),) + + +def test_expectations_differing_only_in_conditions_compare_unequal(): + assert ScoringExpectation(objective="a") != ScoringExpectation(objective="a", conditions=(MatchesObjective(),)) + + +def test_matches_objective_carries_no_text_of_its_own(): + assert dataclasses.fields(MatchesObjective()) == () + + +def test_matches_objective_is_a_condition(): + assert isinstance(MatchesObjective(), Condition) + + +def test_matches_objective_instances_compare_equal(): + assert MatchesObjective() == MatchesObjective() + + +def test_matches_objective_is_frozen(): + condition = MatchesObjective() + + with pytest.raises(dataclasses.FrozenInstanceError): + condition.objective = "something" From ee75dad493cb684a7d58b188dffec7c4eff1223a Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Wed, 19 Aug 2026 14:15:17 -0700 Subject: [PATCH 10/11] Address review of the scorer contract phase 1 Wrapper scorers forwarded a prepared message by re-describing it as a scorable, so children reloaded the pre-substitution pieces from memory, or lost role and error state for a message that was never persisted. Add MessageScorer.score_message_async as the in-hand entry point and route the composite, inverter, threshold, and deprecated-message paths through it. Also: adapt pre-2.0 direct Scorer subclasses so they keep working; declare and enforce routable conditions instead of dropping them silently; read the request for the scored turn rather than the conversation's latest; export MessageScoringOptions and MessageScorableResolver; validate MessageScorable id tuples; migrate the remaining documentation call sites. --- doc/code/memory/5_advanced_memory.ipynb | 9 +- doc/code/memory/5_advanced_memory.py | 9 +- doc/code/scoring/1_true_false_scorers.ipynb | 10 +- doc/code/scoring/1_true_false_scorers.py | 10 +- doc/code/scoring/2_float_scale_scorers.ipynb | 8 +- doc/code/scoring/2_float_scale_scorers.py | 8 +- doc/code/scoring/3_combining_scorers.ipynb | 4 +- doc/code/scoring/3_combining_scorers.py | 4 +- doc/code/targets/round_robin_target.ipynb | 4 +- doc/code/targets/round_robin_target.py | 4 +- pyrit/executor/attack/multi_turn/crescendo.py | 2 +- .../executor/attack/multi_turn/red_teaming.py | 3 +- pyrit/models/score/scorable.py | 18 ++- pyrit/score/__init__.py | 5 +- pyrit/score/message_scorer.py | 146 ++++++++++++++---- pyrit/score/scorer.py | 112 +++++++++++++- .../float_scale_threshold_scorer.py | 20 ++- .../true_false/true_false_composite_scorer.py | 22 ++- .../true_false/true_false_inverter_scorer.py | 20 ++- .../attack/multi_turn/test_tree_of_attacks.py | 2 +- tests/unit/models/test_scorable.py | 19 +++ .../test_float_scale_threshold_scorer.py | 8 +- tests/unit/score/test_message_scorer.py | 134 +++++++++++++++- tests/unit/score/test_scorer.py | 65 ++++++++ 24 files changed, 547 insertions(+), 99 deletions(-) diff --git a/doc/code/memory/5_advanced_memory.ipynb b/doc/code/memory/5_advanced_memory.ipynb index 9a8f6b257b..7cefe0b502 100644 --- a/doc/code/memory/5_advanced_memory.ipynb +++ b/doc/code/memory/5_advanced_memory.ipynb @@ -602,7 +602,7 @@ } ], "source": [ - "from pyrit.models import Message\n", + "from pyrit.models import Message, MessageScorable\n", "from pyrit.score import SubStringScorer\n", "\n", "# Create three scorers with different substrings\n", @@ -623,9 +623,10 @@ "\n", "# Score every response with both scorers — scores are automatically persisted in memory\n", "for msg in assistant_messages:\n", - " await scorer_molotov.score_async(msg) # type: ignore\n", - " await scorer_launder.score_async(msg) # type: ignore\n", - " await scorer_assist.score_async(msg) # type: ignore\n", + " scorable = MessageScorable.from_message(msg)\n", + " await scorer_molotov.score_async(scorable=scorable) # type: ignore\n", + " await scorer_launder.score_async(scorable=scorable) # type: ignore\n", + " await scorer_assist.score_async(scorable=scorable) # type: ignore\n", "\n", "print(f\"Scored {len(assistant_messages)} messages with all three scorers.\")" ] diff --git a/doc/code/memory/5_advanced_memory.py b/doc/code/memory/5_advanced_memory.py index 366bf80bf5..a8e882100a 100644 --- a/doc/code/memory/5_advanced_memory.py +++ b/doc/code/memory/5_advanced_memory.py @@ -246,7 +246,7 @@ # retrieve only the scores produced by a specific scorer. # %% -from pyrit.models import Message +from pyrit.models import Message, MessageScorable from pyrit.score import SubStringScorer # Create three scorers with different substrings @@ -267,9 +267,10 @@ # Score every response with both scorers — scores are automatically persisted in memory for msg in assistant_messages: - await scorer_molotov.score_async(msg) # type: ignore - await scorer_launder.score_async(msg) # type: ignore - await scorer_assist.score_async(msg) # type: ignore + scorable = MessageScorable.from_message(msg) + await scorer_molotov.score_async(scorable=scorable) # type: ignore + await scorer_launder.score_async(scorable=scorable) # type: ignore + await scorer_assist.score_async(scorable=scorable) # type: ignore print(f"Scored {len(assistant_messages)} messages with all three scorers.") diff --git a/doc/code/scoring/1_true_false_scorers.ipynb b/doc/code/scoring/1_true_false_scorers.ipynb index a312e30dd1..89dd688228 100644 --- a/doc/code/scoring/1_true_false_scorers.ipynb +++ b/doc/code/scoring/1_true_false_scorers.ipynb @@ -314,7 +314,7 @@ } ], "source": [ - "from pyrit.models import MessagePiece\n", + "from pyrit.models import MessagePiece, ScoringExpectation\n", "from pyrit.score import SelfAskRefusalScorer\n", "\n", "refusal_scorer = SelfAskRefusalScorer(chat_target=OpenAIChatTarget())\n", @@ -331,8 +331,12 @@ "refusal = MessagePiece(role=\"assistant\", original_value=\"Sorry, but I can't help with that.\").to_message()\n", "refusal.set_response_not_in_memory()\n", "\n", - "blocked_score = (await refusal_scorer.score_async(message=blocked))[0] # type: ignore\n", - "refusal_score = (await refusal_scorer.score_async(message=refusal, objective=\"Explain how to pick a lock\"))[0] # type: ignore\n", + "blocked_score = (await refusal_scorer.score_message_async(message=blocked))[0] # type: ignore\n", + "refusal_score = (\n", + " await refusal_scorer.score_message_async(\n", + " message=refusal, expectation=ScoringExpectation(objective=\"Explain how to pick a lock\")\n", + " )\n", + ")[0] # type: ignore\n", "\n", "print(f\"[refusal] blocked response -> {blocked_score.get_value()}\")\n", "print(f\"[refusal] explicit refusal -> {refusal_score.get_value()}\")" diff --git a/doc/code/scoring/1_true_false_scorers.py b/doc/code/scoring/1_true_false_scorers.py index 12529269b3..e5950f9391 100644 --- a/doc/code/scoring/1_true_false_scorers.py +++ b/doc/code/scoring/1_true_false_scorers.py @@ -149,7 +149,7 @@ # which is common with content filters. A partially blocked response still carries content, so # it is scored normally by the LLM rather than short-circuited. # %% -from pyrit.models import MessagePiece +from pyrit.models import MessagePiece, ScoringExpectation from pyrit.score import SelfAskRefusalScorer refusal_scorer = SelfAskRefusalScorer(chat_target=OpenAIChatTarget()) @@ -166,8 +166,12 @@ refusal = MessagePiece(role="assistant", original_value="Sorry, but I can't help with that.").to_message() refusal.set_response_not_in_memory() -blocked_score = (await refusal_scorer.score_async(message=blocked))[0] # type: ignore -refusal_score = (await refusal_scorer.score_async(message=refusal, objective="Explain how to pick a lock"))[0] # type: ignore +blocked_score = (await refusal_scorer.score_message_async(message=blocked))[0] # type: ignore +refusal_score = ( + await refusal_scorer.score_message_async( + message=refusal, expectation=ScoringExpectation(objective="Explain how to pick a lock") + ) +)[0] # type: ignore print(f"[refusal] blocked response -> {blocked_score.get_value()}") print(f"[refusal] explicit refusal -> {refusal_score.get_value()}") diff --git a/doc/code/scoring/2_float_scale_scorers.ipynb b/doc/code/scoring/2_float_scale_scorers.ipynb index b2713e0c96..014bac45e8 100644 --- a/doc/code/scoring/2_float_scale_scorers.ipynb +++ b/doc/code/scoring/2_float_scale_scorers.ipynb @@ -99,7 +99,7 @@ "\n", "from pyrit.auth import get_azure_token_provider\n", "from pyrit.memory import CentralMemory\n", - "from pyrit.models import Message, MessagePiece\n", + "from pyrit.models import Message, MessagePiece, MessageScorable\n", "from pyrit.score import AzureContentFilterScorer\n", "\n", "azure_content_filter = AzureContentFilterScorer(\n", @@ -120,7 +120,7 @@ "# The score table has a foreign key on the message, so write it to memory first.\n", "CentralMemory.get_memory_instance().add_message_to_memory(request=response)\n", "\n", - "scores = await azure_content_filter.score_async(response) # type: ignore\n", + "scores = await azure_content_filter.score_async(scorable=MessageScorable.from_message(response)) # type: ignore\n", "for score in scores:\n", " # One score per harm category; score_metadata holds the original 0-7 severity.\n", " print(f\"{score.score_category}: value={score.get_value()} metadata={score.score_metadata}\")" @@ -245,7 +245,7 @@ } ], "source": [ - "from pyrit.models import MessagePiece\n", + "from pyrit.models import MessagePiece, MessageScorable\n", "from pyrit.score import InsecureCodeScorer\n", "\n", "insecure_code_scorer = InsecureCodeScorer.from_harm_categories(chat_target=OpenAIChatTarget())\n", @@ -258,7 +258,7 @@ "request = MessagePiece(role=\"assistant\", original_value=snippet, conversation_id=str(uuid4())).to_message()\n", "insecure_code_scorer._memory.add_message_to_memory(request=request)\n", "\n", - "scored = (await insecure_code_scorer.score_async(request))[0] # type: ignore\n", + "scored = (await insecure_code_scorer.score_async(scorable=MessageScorable.from_message(request)))[0] # type: ignore\n", "print(f\"[insecure code] risk={scored.get_value()}\")\n", "print(f\"rationale: {scored.score_rationale}\")" ] diff --git a/doc/code/scoring/2_float_scale_scorers.py b/doc/code/scoring/2_float_scale_scorers.py index f0dca1d94c..7ec6563c37 100644 --- a/doc/code/scoring/2_float_scale_scorers.py +++ b/doc/code/scoring/2_float_scale_scorers.py @@ -43,7 +43,7 @@ from pyrit.auth import get_azure_token_provider from pyrit.memory import CentralMemory -from pyrit.models import Message, MessagePiece +from pyrit.models import Message, MessagePiece, MessageScorable from pyrit.score import AzureContentFilterScorer azure_content_filter = AzureContentFilterScorer( @@ -64,7 +64,7 @@ # The score table has a foreign key on the message, so write it to memory first. CentralMemory.get_memory_instance().add_message_to_memory(request=response) -scores = await azure_content_filter.score_async(response) # type: ignore +scores = await azure_content_filter.score_async(scorable=MessageScorable.from_message(response)) # type: ignore for score in scores: # One score per harm category; score_metadata holds the original 0-7 severity. print(f"{score.score_category}: value={score.get_value()} metadata={score.score_metadata}") @@ -117,7 +117,7 @@ # # Rates how risky a code snippet is, flagging vulnerabilities like injection or weak auth. # %% -from pyrit.models import MessagePiece +from pyrit.models import MessagePiece, MessageScorable from pyrit.score import InsecureCodeScorer insecure_code_scorer = InsecureCodeScorer.from_harm_categories(chat_target=OpenAIChatTarget()) @@ -130,7 +130,7 @@ def authenticate_user(username, password): request = MessagePiece(role="assistant", original_value=snippet, conversation_id=str(uuid4())).to_message() insecure_code_scorer._memory.add_message_to_memory(request=request) -scored = (await insecure_code_scorer.score_async(request))[0] # type: ignore +scored = (await insecure_code_scorer.score_async(scorable=MessageScorable.from_message(request)))[0] # type: ignore print(f"[insecure code] risk={scored.get_value()}") print(f"rationale: {scored.score_rationale}") diff --git a/doc/code/scoring/3_combining_scorers.ipynb b/doc/code/scoring/3_combining_scorers.ipynb index 04fd8385aa..750d84a9b8 100644 --- a/doc/code/scoring/3_combining_scorers.ipynb +++ b/doc/code/scoring/3_combining_scorers.ipynb @@ -296,7 +296,7 @@ "import uuid\n", "\n", "from pyrit.memory import CentralMemory\n", - "from pyrit.models import MessagePiece\n", + "from pyrit.models import MessagePiece, MessageScorable\n", "from pyrit.score import create_conversation_scorer\n", "\n", "memory = CentralMemory.get_memory_instance()\n", @@ -317,7 +317,7 @@ "conversation_scorer = create_conversation_scorer(scorer=persona_breach_scorer)\n", "\n", "# Any message from the conversation works as the trigger.\n", - "score = (await conversation_scorer.score_async(turns[0]))[0] # type: ignore\n", + "score = (await conversation_scorer.score_async(scorable=MessageScorable.from_message(turns[0])))[0] # type: ignore\n", "print(f\"[conversation] persona breach across turns -> {score.get_value()}\")" ] }, diff --git a/doc/code/scoring/3_combining_scorers.py b/doc/code/scoring/3_combining_scorers.py index e665e2b58e..bb7db694f5 100644 --- a/doc/code/scoring/3_combining_scorers.py +++ b/doc/code/scoring/3_combining_scorers.py @@ -160,7 +160,7 @@ import uuid from pyrit.memory import CentralMemory -from pyrit.models import MessagePiece +from pyrit.models import MessagePiece, MessageScorable from pyrit.score import create_conversation_scorer memory = CentralMemory.get_memory_instance() @@ -181,7 +181,7 @@ conversation_scorer = create_conversation_scorer(scorer=persona_breach_scorer) # Any message from the conversation works as the trigger. -score = (await conversation_scorer.score_async(turns[0]))[0] # type: ignore +score = (await conversation_scorer.score_async(scorable=MessageScorable.from_message(turns[0])))[0] # type: ignore print(f"[conversation] persona breach across turns -> {score.get_value()}") # %% [markdown] diff --git a/doc/code/targets/round_robin_target.ipynb b/doc/code/targets/round_robin_target.ipynb index af179710f0..c29f92e530 100644 --- a/doc/code/targets/round_robin_target.ipynb +++ b/doc/code/targets/round_robin_target.ipynb @@ -99,7 +99,7 @@ "import os\n", "\n", "from pyrit.auth import get_azure_openai_auth\n", - "from pyrit.models import Message\n", + "from pyrit.models import Message, MessageScorable\n", "from pyrit.prompt_normalizer import PromptNormalizer\n", "from pyrit.prompt_target import OpenAIChatTarget, RoundRobinTarget\n", "from pyrit.setup import IN_MEMORY, initialize_pyrit_async\n", @@ -788,7 +788,7 @@ "# You may want to use `score_prompts_batch_async` like below in practice for efficiency\n", "# await scorer.score_prompts_batch_async(messages=response_messages) # type: ignore\n", "for i, response_message in enumerate(response_messages):\n", - " scores = await scorer.score_async(message=response_message) # type: ignore\n", + " scores = await scorer.score_async(scorable=MessageScorable.from_message(response_message)) # type: ignore\n", "\n", " # The scorer's internal LLM response has inner_target_identifier in metadata.\n", " # We can check the round-robin counter to determine which target was used.\n", diff --git a/doc/code/targets/round_robin_target.py b/doc/code/targets/round_robin_target.py index ba17eb84d8..806a456af4 100644 --- a/doc/code/targets/round_robin_target.py +++ b/doc/code/targets/round_robin_target.py @@ -38,7 +38,7 @@ import os from pyrit.auth import get_azure_openai_auth -from pyrit.models import Message +from pyrit.models import Message, MessageScorable from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import OpenAIChatTarget, RoundRobinTarget from pyrit.setup import IN_MEMORY, initialize_pyrit_async @@ -252,7 +252,7 @@ # You may want to use `score_prompts_batch_async` like below in practice for efficiency # await scorer.score_prompts_batch_async(messages=response_messages) # type: ignore for i, response_message in enumerate(response_messages): - scores = await scorer.score_async(message=response_message) # type: ignore + scores = await scorer.score_async(scorable=MessageScorable.from_message(response_message)) # type: ignore # The scorer's internal LLM response has inner_target_identifier in metadata. # We can check the round-robin counter to determine which target was used. diff --git a/pyrit/executor/attack/multi_turn/crescendo.py b/pyrit/executor/attack/multi_turn/crescendo.py index 0d9ce56cb7..371d3b9076 100644 --- a/pyrit/executor/attack/multi_turn/crescendo.py +++ b/pyrit/executor/attack/multi_turn/crescendo.py @@ -40,12 +40,12 @@ from pyrit.score import ( FloatScaleThresholdScorer, MessageScorable, + MessageScoringOptions, NumericRubric, Scorer, SelfAskRefusalScorer, SelfAskScaleScorer, ) -from pyrit.score.message_scorer import MessageScoringOptions from pyrit.score.score_utils import normalize_score_to_float if TYPE_CHECKING: diff --git a/pyrit/executor/attack/multi_turn/red_teaming.py b/pyrit/executor/attack/multi_turn/red_teaming.py index 06af753d31..c9cb196f36 100644 --- a/pyrit/executor/attack/multi_turn/red_teaming.py +++ b/pyrit/executor/attack/multi_turn/red_teaming.py @@ -39,8 +39,7 @@ from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import CapabilityName from pyrit.prompt_target.common.target_requirements import TargetRequirements -from pyrit.score import MessageScorable -from pyrit.score.message_scorer import MessageScoringOptions +from pyrit.score import MessageScorable, MessageScoringOptions if TYPE_CHECKING: from collections.abc import Callable diff --git a/pyrit/models/score/scorable.py b/pyrit/models/score/scorable.py index 9185a3179a..27c5eeebd5 100644 --- a/pyrit/models/score/scorable.py +++ b/pyrit/models/score/scorable.py @@ -35,6 +35,19 @@ class MessageScorable(Scorable): message_piece_ids: tuple[uuid.UUID | str, ...] + def __post_init__(self) -> None: + """ + Reject id tuples that cannot name evidence. + + Raises: + ValueError: If no ids are given, or if an id is repeated. + """ + if not self.message_piece_ids: + raise ValueError("A MessageScorable must name at least one message piece.") + seen = [str(piece_id) for piece_id in self.message_piece_ids] + if len(set(seen)) != len(seen): + raise ValueError(f"A MessageScorable must name each message piece once, got {seen}.") + @classmethod def from_message( cls, @@ -70,7 +83,10 @@ def from_message(cls, message: Message) -> ContentScorable: Describe the converted content of a single-piece ephemeral message. Scorers consume ``converted_value``, so this adapter preserves the converted value - and data type rather than the pre-conversion input. + and data type rather than the pre-conversion input. Everything else the message + carried is dropped, including its role and its error state, so a scorer's + deterministic blocked-response handling no longer applies. Use + ``MessageScorer.score_message_async`` when that state is part of the evidence. Args: message (Message): The ephemeral message whose converted content to take. diff --git a/pyrit/score/__init__.py b/pyrit/score/__init__.py index 79460fb5b0..78397b6023 100644 --- a/pyrit/score/__init__.py +++ b/pyrit/score/__init__.py @@ -31,7 +31,8 @@ render_likert_system_prompt, ) from pyrit.score.float_scale.self_ask_scale_scorer import SelfAskScaleScorer, render_scale_system_prompt -from pyrit.score.message_scorer import MessageScorer +from pyrit.score.message_scorable_resolver import MessageScorableResolver +from pyrit.score.message_scorer import MessageScorer, MessageScoringOptions from pyrit.score.response_handler import CallableResponseHandler, JsonSchemaResponseHandler, ResponseHandler from pyrit.score.scorable import ContentScorable, MessageScorable, Scorable from pyrit.score.scorer import Scorer @@ -191,7 +192,9 @@ def __getattr__(name: str) -> object: "LlamaGuardScorer", "MarkdownInjectionScorer", "MessageScorable", + "MessageScorableResolver", "MessageScorer", + "MessageScoringOptions", "MethKeywordScorer", "MetricsType", "NerveAgentKeywordScorer", diff --git a/pyrit/score/message_scorer.py b/pyrit/score/message_scorer.py index 9e454d8d7b..12cc0d518f 100644 --- a/pyrit/score/message_scorer.py +++ b/pyrit/score/message_scorer.py @@ -7,16 +7,16 @@ import logging from abc import abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, cast +from typing import TYPE_CHECKING, ClassVar, cast from pyrit.common.deprecation import print_deprecation_message from pyrit.exceptions import PyritException, ScorerLLMResponseBlockedException from pyrit.models import ( ChatMessageRole, - ContentScorable, + Condition, + MatchesObjective, Message, MessagePiece, - MessageScorable, PromptResponseError, Scorable, Score, @@ -62,22 +62,24 @@ def extract_objective_from_previous_turn(*, message: Message, memory: MemoryInte if not message.message_pieces: return "" - piece = message.get_piece() + scored_piece = message.get_piece() - if piece.api_role != "assistant": + if scored_piece.api_role != "assistant": return "" - conversation = memory.get_message_pieces(conversation_id=piece.conversation_id) - if not conversation: + # The request is the turn before the response being scored, not before whatever the + # conversation has grown to since. Scoring an earlier response must not read the latest turn. + previous_sequence = scored_piece.sequence - 1 + if previous_sequence < 0: return "" - last_prompt = max(conversation, key=lambda x: x.sequence) + conversation = memory.get_message_pieces(conversation_id=scored_piece.conversation_id) return "\n".join( [ piece.original_value for piece in conversation - if piece.sequence == last_prompt.sequence - 1 and piece.original_value_data_type == "text" + if piece.sequence == previous_sequence and piece.original_value_data_type == "text" ] ) @@ -103,6 +105,11 @@ class MessageScorer(Scorer): #: family's neutral fallback score instead of raising. raise_if_scorer_blocks: bool = True + #: A message scorer judges the message against the expectation's objective, so it + #: routes ``MatchesObjective``. Domain-specific conditions belong to the scorers + #: that match them. + SUPPORTED_CONDITIONS: ClassVar[frozenset[type[Condition]]] = frozenset({MatchesObjective}) + def __init__( self, *, @@ -122,6 +129,24 @@ def __init__( self._message_resolver = message_resolver or MessageScorableResolver() super().__init__(chat_target=chat_target) + def _validate_expectation(self, *, expectation: ScoringExpectation | None) -> None: + """ + Reject conditions this scorer cannot consume, and unusable ``MatchesObjective``. + + Raises: + ValueError: If ``MatchesObjective`` is present without an objective to match. + """ + super()._validate_expectation(expectation=expectation) + if expectation is None: + return + if any(isinstance(condition, MatchesObjective) for condition in expectation.conditions) and not ( + expectation.objective + ): + raise ValueError( + "MatchesObjective requires the expectation to carry an objective. " + "Set ScoringExpectation.objective or drop the condition." + ) + async def score_async( self, message: Message | None = None, @@ -150,7 +175,7 @@ async def score_async( Returns: list[Score]: The persisted scores, or an empty list when policy skips the message. """ - resolved_scorable, resolved_expectation, options, infer_objective = self._consolidate_message_inputs( + resolved_expectation, options, infer_objective = self._consolidate_message_inputs( message=message, scorable=scorable, expectation=expectation, @@ -160,11 +185,56 @@ async def score_async( skip_on_error_result=skip_on_error_result, infer_objective_from_request=infer_objective_from_request, ) - scores = await self._score_message_scorable_async( - scorable=resolved_scorable, - expectation=resolved_expectation, - options=options, - infer_objective_from_request=infer_objective, + self._validate_expectation(expectation=resolved_expectation) + + # The deprecated parameter hands over the message itself, so scoring it must not round + # trip through a reference. Re-describing it would reload the persisted originals and + # drop role and error state for a message that was never persisted at all. + if message is not None: + scores = await self._score_resolved_message_async( + message=message, + expectation=resolved_expectation, + options=options, + infer_objective_from_request=infer_objective, + ) + else: + scores = await self._score_message_scorable_async( + scorable=cast("Scorable", scorable), + expectation=resolved_expectation, + options=options, + infer_objective_from_request=infer_objective, + ) + return self._validate_and_persist_scores(scores=scores) + + async def score_message_async( + self, + *, + message: Message, + expectation: ScoringExpectation | None = None, + message_options: MessageScoringOptions | None = None, + ) -> list[Score]: + """ + Score a message that is already in hand. + + Use this when the caller holds the message itself rather than a reference to it: + an ephemeral response that was never persisted, or a scoring view a wrapping scorer + has already prepared. Naming persisted evidence with a ``MessageScorable`` stays the + default, because a reference is what a stored score can be audited against. + + Args: + message (Message): The message to score. + expectation (ScoringExpectation | None): What to look for. Defaults to None. + message_options (MessageScoringOptions | None): Message-family policy. Defaults to None. + + Returns: + list[Score]: The persisted scores, or an empty list when policy skips the message. + """ + self._validate_expectation(expectation=expectation) + scores = await self._score_resolved_message_async( + message=message, + expectation=expectation, + options=message_options or MessageScoringOptions(), + infer_objective_from_request=False, ) return self._validate_and_persist_scores(scores=scores) @@ -179,7 +249,7 @@ def _consolidate_message_inputs( role_filter: ChatMessageRole | None, skip_on_error_result: bool | None, infer_objective_from_request: bool | None, - ) -> tuple[Scorable, ScoringExpectation | None, MessageScoringOptions, bool]: + ) -> tuple[ScoringExpectation | None, MessageScoringOptions, bool]: if message is not None and scorable is not None: raise ValueError("Pass either 'message' or 'scorable', not both.") if message is None and scorable is None: @@ -204,21 +274,12 @@ def _consolidate_message_inputs( removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, ) - if scorable is not None: - resolved_scorable = scorable - else: - legacy_message = cast("Message", message) - resolved_scorable = ( - ContentScorable.from_message(legacy_message) - if len(legacy_message.message_pieces) == 1 and legacy_message.get_piece().not_in_memory - else MessageScorable.from_message(legacy_message) - ) resolved_expectation = ScoringExpectation(objective=objective) if objective is not None else expectation options = message_options or MessageScoringOptions( role_filter=role_filter, skip_on_error_result=skip_on_error_result or False, ) - return resolved_scorable, resolved_expectation, options, bool(infer_objective_from_request) + return resolved_expectation, options, bool(infer_objective_from_request) async def _score_scorable_async( self, @@ -262,13 +323,42 @@ async def _score_message_scorable_async( Raises: TypeError: If the scorable is not message-shaped. + """ + message = self._message_resolver.resolve(scorable=scorable, memory=self._memory) + return await self._score_resolved_message_async( + message=message, + expectation=expectation, + options=options, + infer_objective_from_request=infer_objective_from_request, + ) + + async def _score_resolved_message_async( + self, + *, + message: Message, + expectation: ScoringExpectation | None, + options: MessageScoringOptions, + infer_objective_from_request: bool, + ) -> list[Score]: + """ + Run the message-scoring pipeline over an acquired message. + + Args: + message (Message): The acquired message. + expectation (ScoringExpectation | None): What to look for. + options (MessageScoringOptions): Message-only scoring policy. + infer_objective_from_request (bool): Deprecated; read the objective from the + previous turn when the expectation carries none. + + Returns: + list[Score]: The scores, or an empty list when a filter skipped the message. + + Raises: ScorerLLMResponseBlockedException: If the scorer's own LLM response is blocked by content filtering and ``raise_if_scorer_blocks`` is True (the default). PyritException: If scoring raises a PyRIT exception (re-raised with enhanced context). RuntimeError: If scoring raises a non-PyRIT exception (wrapped with scorer context). """ - message = self._message_resolver.resolve(scorable=scorable, memory=self._memory) - objective = expectation.objective if expectation else None # Structured refusals are persisted as blocked error pieces, but scorers should diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 0654006f1b..3cd3a809df 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -14,6 +14,7 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, + Condition, ContentScorable, Identifiable, Message, @@ -35,6 +36,7 @@ from pyrit.score.scorer_evaluation.metrics_type import RegistryUpdateBehavior from pyrit.score.scorer_evaluation.scorer_evaluator import ScorerEvalDatasetFiles from pyrit.score.scorer_evaluation.scorer_metrics import ScorerMetrics + from pyrit.score.scorer_prompt_validator import ScorerPromptValidator logger = logging.getLogger(__name__) @@ -42,6 +44,52 @@ LEGACY_SCORE_ASYNC_REMOVED_IN = "2.0.0" +async def _legacy_score_scorable_async( + self: Scorer, + *, + scorable: Scorable, + expectation: ScoringExpectation | None, +) -> list[Score]: + """ + Route a scorable to a pre-2.0 subclass that only implements ``_score_async``. + + Returns: + list[Score]: The scores the legacy scorer body produced. + """ + from pyrit.score.message_scorable_resolver import MessageScorableResolver + + print_deprecation_message( + old_item=f"{type(self).__name__}._score_async on a direct Scorer subclass", + new_item="pyrit.score.MessageScorer (or TrueFalseScorer / FloatScaleScorer) as the base class", + removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, + ) + resolver = getattr(self, "_message_resolver", None) or MessageScorableResolver() + message = resolver.resolve(scorable=scorable, memory=self._memory) + legacy_score_async = self._score_async # type: ignore[ty:unresolved-attribute] + return await legacy_score_async(message, objective=expectation.objective if expectation else None) + + +def _adapt_legacy_message_scorer(cls: type) -> None: + """ + Give a pre-2.0 direct ``Scorer`` subclass an implementation of the scorable contract. + + Subclasses of ``MessageScorer`` already inherit one, so they are left alone. A class that + predates the split implements ``_score_async`` instead, and would otherwise fail to + instantiate because ``_score_scorable_async`` is abstract. ``ABCMeta`` recomputes + ``__abstractmethods__`` after ``__init_subclass__``, so assigning it here is enough. + """ + for base in cls.__mro__: + if base is Scorer: + break + if "_score_scorable_async" in base.__dict__: + return + + if not any("_score_async" in base.__dict__ for base in cls.__mro__): + return + + cls._score_scorable_async = _legacy_score_scorable_async # type: ignore[ty:invalid-assignment] + + class Scorer(Identifiable, abc.ABC): """ Abstract base class for scorers. @@ -62,6 +110,11 @@ class Scorer(Identifiable, abc.ABC): #: validate it. TARGET_REQUIREMENTS: ClassVar[TargetRequirements] = TargetRequirements() + #: Condition types this scorer consumes. An expectation is a routing envelope, so a + #: condition that no scorer in the tree declares is a configuration error rather than + #: something to drop silently. Wrapping scorers report their children's union. + SUPPORTED_CONDITIONS: ClassVar[frozenset[type[Condition]]] = frozenset() + _identifier: ComponentIdentifier | None = None def __init_subclass__(cls, **kwargs: Any) -> None: @@ -75,18 +128,43 @@ def __init_subclass__(cls, **kwargs: Any) -> None: from pyrit.common.brick_contract import enforce_keyword_only_init enforce_keyword_only_init(cls, base_name="Scorer") + _adapt_legacy_message_scorer(cls) - def __init__(self, *, chat_target: PromptTarget | None = None) -> None: + def __init__( + self, + *, + chat_target: PromptTarget | None = None, + validator: ScorerPromptValidator | None = None, + ) -> None: """ Initialize the Scorer. Args: chat_target (PromptTarget | None): Chat target used by the scorer, if any. When provided, it is validated against ``TARGET_REQUIREMENTS``. - """ + validator (ScorerPromptValidator | None): Deprecated. Message validation moved to + ``MessageScorer``; a value passed here is kept so pre-2.0 subclasses keep working. + """ + if validator is not None: + print_deprecation_message( + old_item="Scorer.__init__(validator=...)", + new_item="MessageScorer.__init__(validator=...)", + removed_in=LEGACY_SCORE_ASYNC_REMOVED_IN, + ) + if getattr(self, "_validator", None) is None: + self._validator = validator if chat_target is not None: type(self).TARGET_REQUIREMENTS.validate(target=chat_target) + def supported_conditions(self) -> frozenset[type[Condition]]: + """ + Return the condition types this scorer can consume. + + Returns: + frozenset[type[Condition]]: The declared condition types. + """ + return type(self).SUPPORTED_CONDITIONS + def get_chat_target(self) -> PromptTarget | None: """ Return the chat target used by this scorer, or None if it doesn't use one. @@ -199,9 +277,39 @@ async def score_async( Raises: TypeError: If this scorer does not support this kind of scorable. """ + self._validate_expectation(expectation=expectation) scores = await self._score_scorable_async(scorable=scorable, expectation=expectation) return self._validate_and_persist_scores(scores=scores) + def _validate_expectation(self, *, expectation: ScoringExpectation | None) -> None: + """ + Reject conditions no scorer in this tree consumes, and ambiguous routing. + + Raises: + ValueError: If a condition is unsupported, or if more than one condition of the + same supported type is present. + """ + if expectation is None or not expectation.conditions: + return + + supported = self.supported_conditions() + unsupported = [condition for condition in expectation.conditions if not isinstance(condition, tuple(supported))] + if unsupported: + names = ", ".join(sorted({type(condition).__name__ for condition in unsupported})) + supported_names = ", ".join(sorted(cls.__name__ for cls in supported)) or "none" + raise ValueError( + f"{type(self).__name__} does not consume the condition(s) {names}. " + f"Supported conditions: {supported_names}." + ) + + for condition_type in supported: + matches = [condition for condition in expectation.conditions if isinstance(condition, condition_type)] + if len(matches) > 1: + raise ValueError( + f"{type(self).__name__} received {len(matches)} {condition_type.__name__} conditions. " + "A scorer consumes at most one condition of a given type." + ) + def _validate_and_persist_scores(self, *, scores: list[Score]) -> list[Score]: """ Validate and persist non-empty scorer output. diff --git a/pyrit/score/true_false/float_scale_threshold_scorer.py b/pyrit/score/true_false/float_scale_threshold_scorer.py index 74e5bc95a6..7598dbd678 100644 --- a/pyrit/score/true_false/float_scale_threshold_scorer.py +++ b/pyrit/score/true_false/float_scale_threshold_scorer.py @@ -10,10 +10,9 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, - ContentScorable, + Condition, Message, MessagePiece, - MessageScorable, Score, ScoringExpectation, ) @@ -88,6 +87,15 @@ def get_chat_target(self) -> "PromptTarget | None": """ return self._scorer.get_chat_target() + def supported_conditions(self) -> frozenset[type[Condition]]: + """ + Report what the wrapped scorer consumes. + + Returns: + frozenset[type[Condition]]: The condition types the wrapped scorer routes. + """ + return self._scorer.supported_conditions() + async def _score_async( self, message: Message, @@ -107,12 +115,8 @@ async def _score_async( Returns: list[Score]: A list containing a single true/false Score object based on the threshold comparison. """ - scores = await self._scorer.score_async( - scorable=( - ContentScorable.from_message(message) - if len(message.message_pieces) == 1 and message.get_piece().not_in_memory - else MessageScorable.from_message(message) - ), + scores = await self._scorer.score_message_async( + message=message, expectation=ScoringExpectation(objective=objective), ) diff --git a/pyrit/score/true_false/true_false_composite_scorer.py b/pyrit/score/true_false/true_false_composite_scorer.py index 6236abd748..c4ab9dd751 100644 --- a/pyrit/score/true_false/true_false_composite_scorer.py +++ b/pyrit/score/true_false/true_false_composite_scorer.py @@ -10,10 +10,9 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, - ContentScorable, + Condition, Message, MessagePiece, - MessageScorable, Score, ScoringExpectation, ) @@ -84,6 +83,15 @@ def get_chat_target(self) -> "PromptTarget | None": return target return None + def supported_conditions(self) -> frozenset[type[Condition]]: + """ + Report the union of what the constituent scorers consume. + + Returns: + frozenset[type[Condition]]: The condition types this composite routes. + """ + return frozenset().union(*(scorer.supported_conditions() for scorer in self._scorers)) + async def _score_async( self, message: Message, @@ -106,13 +114,11 @@ async def _score_async( ValueError: If any constituent scorer does not return exactly one score. ValueError: If no scores are generated from the request response pieces. """ - scorable = ( - ContentScorable.from_message(message) - if len(message.message_pieces) == 1 and message.get_piece().not_in_memory - else MessageScorable.from_message(message) - ) + # The children score the evidence this scorer was handed, substitutions and all. + # Naming it instead would send them back to memory for the pre-substitution pieces, + # or discard the role and error state of a message that was never persisted. expectation = ScoringExpectation(objective=objective) - tasks = [scorer.score_async(scorable=scorable, expectation=expectation) for scorer in self._scorers] + tasks = [scorer.score_message_async(message=message, expectation=expectation) for scorer in self._scorers] # Run all response scorings concurrently score_list_results = await asyncio.gather(*tasks) diff --git a/pyrit/score/true_false/true_false_inverter_scorer.py b/pyrit/score/true_false/true_false_inverter_scorer.py index 0fdec6652c..4093a50c2a 100644 --- a/pyrit/score/true_false/true_false_inverter_scorer.py +++ b/pyrit/score/true_false/true_false_inverter_scorer.py @@ -10,10 +10,9 @@ from pyrit.models import ( ChatMessageRole, ComponentIdentifier, - ContentScorable, + Condition, Message, MessagePiece, - MessageScorable, Score, ScoringExpectation, ) @@ -63,6 +62,15 @@ def get_chat_target(self) -> "PromptTarget | None": """ return self._scorer.get_chat_target() + def supported_conditions(self) -> frozenset[type[Condition]]: + """ + Report what the wrapped scorer consumes. + + Returns: + frozenset[type[Condition]]: The condition types the wrapped scorer routes. + """ + return self._scorer.supported_conditions() + async def _score_async( self, message: Message, @@ -82,12 +90,8 @@ async def _score_async( Returns: list[Score]: A list containing a single Score object with the inverted true/false value. """ - scores = await self._scorer.score_async( - scorable=( - ContentScorable.from_message(message) - if len(message.message_pieces) == 1 and message.get_piece().not_in_memory - else MessageScorable.from_message(message) - ), + scores = await self._scorer.score_message_async( + message=message, expectation=ScoringExpectation(objective=objective), ) inv_score = scores[0] diff --git a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py index 4fc630bb99..246bf8975c 100644 --- a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py +++ b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py @@ -367,7 +367,7 @@ async def create_threshold_score_async(*, original_float_value: float, threshold class_module="test_module", ), ) - mock_float_scorer.score_async = AsyncMock(return_value=[float_score]) + mock_float_scorer.score_message_async = AsyncMock(return_value=[float_score]) # Create the actual FloatScaleThresholdScorer threshold_scorer = FloatScaleThresholdScorer(scorer=mock_float_scorer, threshold=threshold) diff --git a/tests/unit/models/test_scorable.py b/tests/unit/models/test_scorable.py index 006a7f5ec2..a4dbb4c032 100644 --- a/tests/unit/models/test_scorable.py +++ b/tests/unit/models/test_scorable.py @@ -67,6 +67,25 @@ def test_message_scorable_from_message_names_pieces(): assert not hasattr(scorable, "message") +def test_message_scorable_rejects_empty_ids(): + with pytest.raises(ValueError, match="at least one message piece"): + MessageScorable(message_piece_ids=()) + + +def test_message_scorable_rejects_duplicate_ids(): + piece_id = uuid.uuid4() + + with pytest.raises(ValueError, match="each message piece once"): + MessageScorable(message_piece_ids=(piece_id, piece_id)) + + +def test_message_scorable_rejects_ids_that_repeat_across_types(): + piece_id = uuid.uuid4() + + with pytest.raises(ValueError, match="each message piece once"): + MessageScorable(message_piece_ids=(piece_id, str(piece_id))) + + def test_content_scorable_defaults_to_text(): assert ContentScorable(value="hello").data_type == "text" diff --git a/tests/unit/score/test_float_scale_threshold_scorer.py b/tests/unit/score/test_float_scale_threshold_scorer.py index f3449e0c1a..aecf49475a 100644 --- a/tests/unit/score/test_float_scale_threshold_scorer.py +++ b/tests/unit/score/test_float_scale_threshold_scorer.py @@ -21,7 +21,7 @@ def create_mock_float_scorer(score_value: float): class_module="test.mock", ) scorer = AsyncMock() - scorer.score_async = AsyncMock( + scorer.score_message_async = AsyncMock( return_value=[ Score( score_value=str(score_value), @@ -75,7 +75,7 @@ async def test_float_scale_threshold_scorer_returns_single_score_with_multi_cate # Mock a scorer that returns multiple scores (like AzureContentFilterScorer) scorer = AsyncMock() prompt_id = uuid.uuid4() - scorer.score_async = AsyncMock( + scorer.score_message_async = AsyncMock( return_value=[ Score( score_value="0.2", @@ -142,7 +142,7 @@ async def test_float_scale_threshold_scorer_handles_empty_scores(): # Mock a scorer that returns empty list (all pieces filtered) scorer = AsyncMock() - scorer.score_async = AsyncMock(return_value=[]) + scorer.score_message_async = AsyncMock(return_value=[]) # get_identifier() returns a ComponentIdentifier mock_identifier = ComponentIdentifier( class_name="MockScorer", @@ -177,7 +177,7 @@ async def test_float_scale_threshold_scorer_with_raise_on_empty_aggregator(): # Mock a scorer that returns empty list (all pieces filtered) scorer = AsyncMock() - scorer.score_async = AsyncMock(return_value=[]) + scorer.score_message_async = AsyncMock(return_value=[]) # get_identifier() returns a ComponentIdentifier mock_identifier = ComponentIdentifier( class_name="MockScorer", diff --git a/tests/unit/score/test_message_scorer.py b/tests/unit/score/test_message_scorer.py index de6f4a40f7..1eab654560 100644 --- a/tests/unit/score/test_message_scorer.py +++ b/tests/unit/score/test_message_scorer.py @@ -9,7 +9,16 @@ import pytest from pyrit.memory import CentralMemory, MemoryInterface -from pyrit.models import ComponentIdentifier, Message, MessagePiece, Score, ScoringExpectation +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + Condition, + MatchesObjective, + Message, + MessagePiece, + Score, + ScoringExpectation, +) from pyrit.score import ( ContentScorable, MessageScorable, @@ -203,7 +212,10 @@ def test_message_scorer_satisfies_the_scorable_contract(self): assert "_score_piece_async" in MessageScorer.__abstractmethods__ def test_message_dependencies_live_on_message_scorer(self): - assert "validator" not in inspect.signature(Scorer).parameters + # The base keeps 'validator' only as a deprecated shim for pre-2.0 subclasses; the + # dependency itself is required by MessageScorer. + assert inspect.signature(Scorer).parameters["validator"].default is None + assert inspect.signature(MessageScorer).parameters["validator"].default is inspect.Parameter.empty for hook in [ "_build_fallback_score", "_apply_structured_refusal_substitution", @@ -306,19 +318,24 @@ async def test_keyword_message_maps_to_message_scorable(self): assert scorer.scored_messages == [message] - async def test_ephemeral_message_maps_its_converted_view_to_content(self): + async def test_ephemeral_message_keeps_its_own_state(self): + """An in-hand message is scored as it stands, so nothing about it is re-derived.""" scorer = RecordingScorer() message = MessagePiece( - role="user", + role="assistant", original_value="original", converted_value="converted", + response_error="blocked", ).to_message() message.set_response_not_in_memory() with pytest.warns(DeprecationWarning, match="Scorer.score_async"): await scorer.score_async(message) - assert scorer.scored_messages[0].get_value() == "converted" + scored = scorer.scored_messages[0] + assert scored.get_value() == "converted" + assert scored.get_piece().role == "assistant" + assert scored.is_error() async def test_message_does_not_widen_to_the_stored_conversation(self, sqlite_instance: MemoryInterface): """The shim scores the supplied message, never the whole conversation behind it.""" @@ -476,3 +493,110 @@ def test_returns_empty_when_the_conversation_is_not_stored(self, sqlite_instance ).to_message() assert extract_objective_from_previous_turn(message=message, memory=sqlite_instance) == "" + + def test_reads_the_request_for_the_scored_turn_not_the_latest_one(self, sqlite_instance: MemoryInterface): + """Scoring an earlier response must not pick up a request from later in the conversation.""" + conversation_id = str(uuid.uuid4()) + turns: list[tuple[str, ChatMessageRole]] = [ + ("the first request", "user"), + ("the first response", "assistant"), + ("a later request", "user"), + ("a later response", "assistant"), + ] + for value, role in turns: + sqlite_instance.add_message_to_memory( + request=MessagePiece(role=role, original_value=value, conversation_id=conversation_id).to_message() + ) + first_response = sqlite_instance.get_message_pieces(conversation_id=conversation_id)[1].to_message() + + objective = extract_objective_from_previous_turn(message=first_response, memory=sqlite_instance) + + assert objective == "the first request" + + +@pytest.mark.usefixtures("patch_central_database") +class TestInHandMessages: + """A message already in hand is scored as it stands, not re-acquired.""" + + async def test_score_message_async_does_not_read_memory(self): + resolver = MagicMock(spec=MessageScorableResolver) + scorer = RecordingScorer(message_resolver=resolver) + message = _assistant_message("in hand") + + await scorer.score_message_async(message=message) + + resolver.resolve.assert_not_called() + assert scorer.scored_messages == [message] + + async def test_score_message_async_preserves_ephemeral_error_state(self): + scorer = RecordingScorer() + message = MessagePiece( + role="assistant", + original_value="", + original_value_data_type="error", + response_error="blocked", + ).to_message() + message.set_response_not_in_memory() + + await scorer.score_message_async(message=message) + + assert scorer.scored_messages[0].is_error() + + async def test_score_message_async_applies_message_options(self): + scorer = RecordingScorer() + + scores = await scorer.score_message_async( + message=_assistant_message(), + message_options=MessageScoringOptions(role_filter="user"), + ) + + assert scores == [] + + +@pytest.mark.usefixtures("patch_central_database") +class TestConditionRouting: + """An expectation is a routing envelope, so a condition is consumed or refused.""" + + async def test_matches_objective_reaches_a_message_scorer(self): + scorer = RecordingScorer() + + scores = await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + expectation=ScoringExpectation(objective="an objective", conditions=(MatchesObjective(),)), + ) + + assert len(scores) == 1 + + async def test_matches_objective_without_an_objective_raises(self): + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="MatchesObjective requires"): + await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + expectation=ScoringExpectation(conditions=(MatchesObjective(),)), + ) + + async def test_unconsumed_condition_raises_instead_of_being_dropped(self): + @dataclasses.dataclass(frozen=True) + class UnroutedCondition(Condition): + pass + + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="does not consume the condition"): + await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + expectation=ScoringExpectation(conditions=(UnroutedCondition(),)), + ) + + async def test_two_conditions_of_one_type_raise(self): + scorer = RecordingScorer() + + with pytest.raises(ValueError, match="at most one condition"): + await scorer.score_async( + scorable=MessageScorable.from_message(_assistant_message()), + expectation=ScoringExpectation( + objective="an objective", + conditions=(MatchesObjective(), MatchesObjective()), + ), + ) diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index 2ee17ccf26..0f65aa16bd 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -1238,6 +1238,71 @@ async def test_base_scorer_score_async_implementation(patch_central_database): assert len(scores) == 2 +class TestLegacyDirectScorerSubclass: + """Scorers written against the pre-2.0 base keep working behind a deprecation warning.""" + + @staticmethod + def _build_legacy_scorer_class(): + class LegacyScorer(Scorer): + def __init__(self, *, validator: ScorerPromptValidator): + super().__init__(validator=validator) + self.scored_messages: list[Message] = [] + + def _build_identifier(self) -> ComponentIdentifier: + return self._create_identifier() + + async def _score_async(self, message: Message, *, objective: str | None = None) -> list[Score]: + self.scored_messages.append(message) + return [ + Score( + score_value="true", + score_value_description="legacy", + score_type="true_false", + score_category=None, + score_metadata=None, + score_rationale="legacy", + scorer_class_identifier=self.get_identifier(), + message_piece_id=message.get_piece().id, + objective=objective, + ) + ] + + def validate_return_scores(self, scores: list[Score]) -> None: + pass + + def get_scorer_metrics(self): + return None + + return LegacyScorer + + def test_legacy_scorer_is_instantiable(self): + legacy_class = self._build_legacy_scorer_class() + + assert "_score_scorable_async" not in legacy_class.__abstractmethods__ + + def test_legacy_validator_argument_warns(self): + legacy_class = self._build_legacy_scorer_class() + + with pytest.warns(DeprecationWarning, match="Scorer.__init__"): + scorer = legacy_class(validator=DummyValidator()) + + assert scorer._validator is not None + + async def test_legacy_scorer_scores_a_scorable(self, patch_central_database): + legacy_class = self._build_legacy_scorer_class() + with pytest.warns(DeprecationWarning): + scorer = legacy_class(validator=DummyValidator()) + message = store_message( + MessagePiece(role="assistant", original_value="legacy response", conversation_id="legacy").to_message() + ) + + with pytest.warns(DeprecationWarning, match="_score_async"): + scores = await scorer.score_async(scorable=MessageScorable.from_message(message)) + + assert len(scores) == 1 + assert scorer.scored_messages[0].get_value() == "legacy response" + + # Tests for get_identifier and identifier From 8929f52dd28897d30988133e8eb4b3dde8bce8f7 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 20 Aug 2026 09:08:23 -0700 Subject: [PATCH 11/11] Fix scorer routing CI failures Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 7affb573-e048-445a-9152-0b46e997dd89 --- pyrit/score/scorer_prompt_validator.py | 2 +- tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pyrit/score/scorer_prompt_validator.py b/pyrit/score/scorer_prompt_validator.py index 5be9812f46..cfb4c8ed77 100644 --- a/pyrit/score/scorer_prompt_validator.py +++ b/pyrit/score/scorer_prompt_validator.py @@ -69,7 +69,7 @@ def __init__( @property def is_objective_required(self) -> bool: - """Return whether the scorer uses the objective as a required criterion.""" + """Whether the scorer uses the objective as a required criterion.""" return self._is_objective_required def validate(self, message: Message, objective: str | None) -> None: diff --git a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py index 356a8eb132..6e9431ce19 100644 --- a/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py +++ b/tests/unit/executor/attack/multi_turn/test_tree_of_attacks.py @@ -367,7 +367,7 @@ async def create_threshold_score_async(*, original_float_value: float, threshold class_module="test_module", ), ) - mock_float_scorer.score_message_async = AsyncMock(return_value=[float_score]) + mock_float_scorer._score_nested_message_async = AsyncMock(return_value=[float_score]) # Create the actual FloatScaleThresholdScorer threshold_scorer = FloatScaleThresholdScorer(scorer=mock_float_scorer, threshold=threshold)