From 7b0bf088853df5fa740592503b66f0062cff4ec4 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 12 Aug 2026 06:32:00 -0700 Subject: [PATCH 1/2] MAINT: Consolidate OpenAI response handling Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../prompt_target/openai/_response_adapter.py | 245 ++++++++++++++++++ .../openai/openai_chat_target.py | 98 +------ .../openai/openai_completion_target.py | 32 +-- .../openai/openai_response_target.py | 210 +-------------- pyrit/prompt_target/openai/openai_target.py | 10 +- .../target/test_openai_response_adapters.py | 158 +++++++++++ 6 files changed, 419 insertions(+), 334 deletions(-) create mode 100644 pyrit/prompt_target/openai/_response_adapter.py create mode 100644 tests/unit/prompt_target/target/test_openai_response_adapters.py diff --git a/pyrit/prompt_target/openai/_response_adapter.py b/pyrit/prompt_target/openai/_response_adapter.py new file mode 100644 index 0000000000..4b3be7c28f --- /dev/null +++ b/pyrit/prompt_target/openai/_response_adapter.py @@ -0,0 +1,245 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import logging +from typing import Any, Protocol, TypeVar + +from openai.types.chat import ChatCompletion +from openai.types.responses import Response, ResponseOutputText + +from pyrit.exceptions import EmptyResponseException, PyritException +from pyrit.models import MessagePiece, TokenUsage, read_usage_int, read_usage_value +from pyrit.prompt_target.common.chat_completions_response_parser import ( + capture_token_usage, + capture_usage_and_finish_reason, + extract_partial_content, + get_finish_reason, + is_content_filter_response, + validate_chat_completion_response, +) +from pyrit.prompt_target.common.utils import ( + set_response_metadata, + set_token_usage_metadata, + warn_truncated_response, +) +from pyrit.prompt_target.openai.openai_error_handling import _is_content_filter_error + +logger = logging.getLogger(__name__) + +ResponseT = TypeVar("ResponseT", contravariant=True) + + +class OpenAIResponseAdapter(Protocol[ResponseT]): + """The response-format contract used by ``OpenAITarget``.""" + + def is_content_filter(self, *, response: ResponseT) -> bool: + """Return whether the response was blocked by a content filter.""" + ... + + def extract_partial_content(self, *, response: ResponseT) -> str | None: + """Extract content emitted before a content filter stopped generation.""" + ... + + def capture_metadata(self, *, response: ResponseT, pieces: list[MessagePiece]) -> None: + """Copy provider response metadata to PyRIT response pieces.""" + ... + + def validate(self, *, response: ResponseT, is_truncated: bool) -> None: + """Validate the provider response.""" + ... + + def is_truncated(self, *, response: ResponseT) -> bool: + """Return whether generation stopped at the output-token limit.""" + ... + + +class NoOpOpenAIResponseAdapter: + """Default behavior for OpenAI targets without a structured response format.""" + + def is_content_filter(self, *, response: Any) -> bool: + """Return False because this format has no content-filter signal.""" + return False + + def extract_partial_content(self, *, response: Any) -> str | None: + """Return no partial content because this format has no extraction rule.""" + return None + + def capture_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: + """Leave response metadata unchanged.""" + + def validate(self, *, response: Any, is_truncated: bool) -> None: + """Accept the response without format-specific validation.""" + + def is_truncated(self, *, response: Any) -> bool: + """Return False because this format has no truncation signal.""" + return False + + +class ChatCompletionsResponseAdapter: + """Response behavior for the OpenAI Chat Completions wire format.""" + + def is_content_filter(self, *, response: ChatCompletion) -> bool: + """Return whether ``finish_reason`` reports a content filter.""" + return is_content_filter_response(response) + + def extract_partial_content(self, *, response: ChatCompletion) -> str | None: + """ + Extract text emitted before a content filter stopped generation. + + Args: + response (ChatCompletion): The provider response. + + Returns: + str | None: Partial text, if present. + """ + return extract_partial_content(response) + + def capture_metadata(self, *, response: ChatCompletion, pieces: list[MessagePiece]) -> None: + """Capture token usage and the first choice's finish reason.""" + capture_usage_and_finish_reason(pieces=pieces, response=response) + + def validate(self, *, response: ChatCompletion, is_truncated: bool) -> None: + """Validate the response while accepting token-limit truncation.""" + if is_truncated: + warn_truncated_response(signal="finish_reason='length'", limit_parameter="max_completion_tokens") + return + validate_chat_completion_response(response=response) + + def is_truncated(self, *, response: ChatCompletion) -> bool: + """Return whether ``finish_reason`` reports token-limit truncation.""" + return get_finish_reason(response=response) == "length" + + +class CompletionsResponseAdapter(NoOpOpenAIResponseAdapter): + """Response behavior for the legacy OpenAI Completions wire format.""" + + def capture_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: + """Capture call-level usage and each choice's finish reason.""" + capture_token_usage(pieces=pieces, response=response) + + choices = getattr(response, "choices", None) or [] + for index, piece in enumerate(pieces): + choice = choices[index] if index < len(choices) else None + set_response_metadata(pieces=[piece], finish_reason=getattr(choice, "finish_reason", None)) + + +class ResponsesResponseAdapter: + """Response behavior for the OpenAI Responses API wire format.""" + + def is_content_filter(self, *, response: Response) -> bool: + """Return whether the response reports content filtering.""" + error = getattr(response, "error", None) + if error is not None and _is_content_filter_error(response.model_dump()): + return True + + if getattr(response, "status", None) != "incomplete": + return False + incomplete_details = getattr(response, "incomplete_details", None) + return incomplete_details is not None and incomplete_details.reason == "content_filter" + + def extract_partial_content(self, *, response: Response) -> str | None: + """ + Extract text from completed message sections in a filtered response. + + Args: + response (Response): The provider response. + + Returns: + str | None: Partial text, if present. + """ + try: + parts = [ + content_item.text + for section in response.output or [] + if getattr(section, "type", None) == "message" and getattr(section, "status", None) == "completed" + for content_item in getattr(section, "content", None) or [] + if isinstance(content_item, ResponseOutputText) and content_item.text + ] + except (AttributeError, IndexError, TypeError): + return None + return "\n".join(parts) if parts else None + + def capture_metadata(self, *, response: Response, pieces: list[MessagePiece]) -> None: + """Capture token usage, response status, and incomplete reason.""" + if not pieces: + return + + usage = getattr(response, "usage", None) + parsed_usage = token_usage_from_responses(usage) if usage is not None else None + set_token_usage_metadata(pieces=pieces, usage=parsed_usage) + + status = getattr(response, "status", None) + incomplete_details = getattr(response, "incomplete_details", None) + incomplete_reason = getattr(incomplete_details, "reason", None) if incomplete_details else None + set_response_metadata(pieces=pieces, status=status, incomplete_reason=incomplete_reason) + + def validate(self, *, response: Response, is_truncated: bool) -> None: + """ + Validate the response while accepting token-limit truncation. + + Args: + response (Response): The provider response. + is_truncated (bool): Whether the target classified the response as token-limit truncation. + + Raises: + PyritException: If the provider reports an error or unexpected status. + EmptyResponseException: If a completed response has no output. + """ + if response.error is not None and response.error.code != "content_filter": + raise PyritException(message=f"Response error: {response.error.code} - {response.error.message}") + + if is_truncated: + warn_truncated_response( + signal="status='incomplete', reason='max_output_tokens'", + limit_parameter="max_output_tokens", + ) + return + + if response.status != "completed": + raise PyritException(message=f"Unexpected status: {response.status}") + + if not response.output: + logger.error("The response returned no valid output.") + raise EmptyResponseException(message="The response returned an empty response.") + + def is_truncated(self, *, response: Response) -> bool: + """Return whether the response stopped at ``max_output_tokens``.""" + if response.status != "incomplete": + return False + incomplete_details = response.incomplete_details + reason = incomplete_details.reason if incomplete_details else None + return reason == "max_output_tokens" + + +def token_usage_from_responses(usage: Any) -> TokenUsage: + """ + Build a ``TokenUsage`` from a Responses API ``usage`` payload. + + Args: + usage (Any): The Responses API usage object. + + Returns: + TokenUsage: The parsed token usage. + """ + input_details = read_usage_value(source=usage, name="input_tokens_details") + output_details = read_usage_value(source=usage, name="output_tokens_details") + + input_tokens = read_usage_int(source=usage, name="input_tokens") + output_tokens = read_usage_int(source=usage, name="output_tokens") + total_tokens = read_usage_int(source=usage, name="total_tokens") + if total_tokens is None and input_tokens is not None and output_tokens is not None: + total_tokens = input_tokens + output_tokens + + extra: dict[str, int] = {} + cache_write_tokens = read_usage_int(source=input_details, name="cache_write_tokens") + if cache_write_tokens is not None: + extra["cache_write_tokens"] = cache_write_tokens + + return TokenUsage( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=total_tokens, + reasoning_tokens=read_usage_int(source=output_details, name="reasoning_tokens"), + cached_tokens=read_usage_int(source=input_details, name="cached_tokens"), + extra=extra, + ) diff --git a/pyrit/prompt_target/openai/openai_chat_target.py b/pyrit/prompt_target/openai/openai_chat_target.py index ff82e3bc78..46fcfc8ce4 100644 --- a/pyrit/prompt_target/openai/openai_chat_target.py +++ b/pyrit/prompt_target/openai/openai_chat_target.py @@ -27,13 +27,8 @@ ) from pyrit.prompt_target.common.chat_completions_response_parser import ( build_response_pieces_async, - capture_usage_and_finish_reason, detect_response_content, - extract_partial_content, - get_finish_reason, - is_content_filter_response, save_audio_response_async, - validate_chat_completion_response, ) from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.prompt_target.common.target_configuration import TargetConfiguration @@ -42,8 +37,8 @@ limit_requests_per_minute, validate_temperature, validate_top_p, - warn_truncated_response, ) +from pyrit.prompt_target.openai._response_adapter import ChatCompletionsResponseAdapter from pyrit.prompt_target.openai.openai_chat_audio_config import OpenAIChatAudioConfig from pyrit.prompt_target.openai.openai_target import OpenAITarget @@ -93,6 +88,7 @@ class OpenAIChatTarget(OpenAITarget): ), ) ) + _response_adapter = ChatCompletionsResponseAdapter() @forward_init_parameters def __init__( @@ -249,96 +245,6 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me ) return [response] - def _check_content_filter(self, response: Any) -> bool: - """ - Check if a Chat Completions API response has finish_reason=content_filter. - - Args: - response: A ChatCompletion object from the OpenAI SDK. - - Returns: - True if content was filtered, False otherwise. - """ - return is_content_filter_response(response) - - def _extract_partial_content(self, response: Any) -> str | None: - """ - Extract partial content from a Chat Completions response with finish_reason=content_filter. - - When Azure Content Safety triggers mid-generation, the model may have produced partial - text in ``response.choices[0].message.content`` before being cut off. - - Args: - response: A ChatCompletion object from the OpenAI SDK. - - Returns: - The partial text content, or None if no content was generated. - """ - return extract_partial_content(response) - - def _capture_response_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: - """ - Record token usage and ``finish_reason`` from a Chat Completions response. - - Args: - response: A ChatCompletion object from the OpenAI SDK, or the synthetic stand-in used - when the SDK raises on a content filter. - pieces (list[MessagePiece]): The constructed response pieces. - """ - capture_usage_and_finish_reason(pieces=pieces, response=response) - - def _validate_response(self, response: ChatCompletion, request: MessagePiece) -> None: - """ - Validate a Chat Completions API response for errors. - - Checks for: - - Missing choices - - Invalid finish_reason - - At least one valid response type (text content, audio, or tool_calls) - - A ``finish_reason == "length"`` (token-limit truncation) response is treated as valid, with a - warning, so that ``_construct_message_from_response_async`` can preserve any partial content - or fall back to a graceful empty response. Genuinely empty responses (no truncation) are - raised so the retry logic can attempt to get a complete response. Content filter responses - are handled separately by ``_check_content_filter``. - - Args: - response: The ChatCompletion response from OpenAI SDK. - request: The original request MessagePiece. - - Raises: - PyritException: For unexpected response structures or finish reasons. - EmptyResponseException: When the API returns an empty response that was not caused by - token-limit truncation. - """ - # Token-limit truncation is handled before the shared validator, which would otherwise raise - # EmptyResponseException on a validly truncated but empty response. Reasoning models can spend - # the whole budget on hidden reasoning before emitting a visible answer, and a low limit may be - # deliberate, so warn instead of raising and let construction preserve any partial content or - # fall back to a graceful empty response. - if self._is_truncated_response(response): - warn_truncated_response(signal="finish_reason='length'", limit_parameter="max_completion_tokens") - return - - # Genuinely empty responses (no truncation) raise so the retry logic can attempt to get a - # complete response. - validate_chat_completion_response(response=response) - - def _is_truncated_response(self, response: ChatCompletion) -> bool: - """ - Return True if the response was cut off by the token limit. - - The Chat Completions API signals token-limit truncation via ``finish_reason == "length"`` - on the first choice. - - Args: - response: A ChatCompletion response from the OpenAI SDK. - - Returns: - bool: True if the response was truncated at the token limit, False otherwise. - """ - return get_finish_reason(response=response) == "length" - def _detect_response_content(self, message: Any) -> tuple[bool, bool, bool]: """ Detect what content types are present in a ChatCompletion message. diff --git a/pyrit/prompt_target/openai/openai_completion_target.py b/pyrit/prompt_target/openai/openai_completion_target.py index 8d067da123..0a84be4608 100644 --- a/pyrit/prompt_target/openai/openai_completion_target.py +++ b/pyrit/prompt_target/openai/openai_completion_target.py @@ -8,13 +8,11 @@ from pyrit.exceptions.exception_classes import ( pyrit_target_retry, ) -from pyrit.models import ComponentIdentifier, Message, MessagePiece, construct_response_from_request -from pyrit.prompt_target.common.chat_completions_response_parser import ( - capture_token_usage, -) +from pyrit.models import ComponentIdentifier, Message, construct_response_from_request from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.prompt_target.common.target_configuration import TargetConfiguration -from pyrit.prompt_target.common.utils import limit_requests_per_minute, set_response_metadata +from pyrit.prompt_target.common.utils import limit_requests_per_minute +from pyrit.prompt_target.openai._response_adapter import CompletionsResponseAdapter from pyrit.prompt_target.openai.openai_target import OpenAITarget logger = logging.getLogger(__name__) @@ -24,6 +22,7 @@ class OpenAICompletionTarget(OpenAITarget): """A prompt target for OpenAI completion endpoints.""" _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration(capabilities=TargetCapabilities()) + _response_adapter = CompletionsResponseAdapter() @forward_init_parameters def __init__( @@ -158,29 +157,6 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me ) return [response] - def _capture_response_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: - """ - Record token usage and each choice's ``finish_reason`` from a Completion response. - - Usage is per-call and lands on the first piece, as everywhere else. ``finish_reason`` is - per-choice, and this is the one target that maps a piece to each choice, so each piece gets - its own: reading only ``choices[0]`` would report one generation's stop reason for all of - them and hide a content filter that tripped on a later choice. - - Args: - response: A Completion object from the OpenAI SDK, or the synthetic stand-in used when - the SDK raises on a content filter. - pieces (list[MessagePiece]): The constructed response pieces. - """ - # The Completions and Chat Completions APIs report the same ``usage`` schema, so the parser is shared. - capture_token_usage(pieces=pieces, response=response) - - choices = getattr(response, "choices", None) or [] - for index, piece in enumerate(pieces): - # Per piece, not per response: each piece is one choice, so each carries its own stop reason. - choice = choices[index] if index < len(choices) else None - set_response_metadata(pieces=[piece], finish_reason=getattr(choice, "finish_reason", None)) - async def _construct_message_from_response_async(self, response: Any, request: Any) -> Message: """ Construct a Message from a Completion response. diff --git a/pyrit/prompt_target/openai/openai_response_target.py b/pyrit/prompt_target/openai/openai_response_target.py index ee94bca135..956dddbceb 100644 --- a/pyrit/prompt_target/openai/openai_response_target.py +++ b/pyrit/prompt_target/openai/openai_response_target.py @@ -30,22 +30,17 @@ MessagePiece, PromptDataType, PromptResponseError, - TokenUsage, - read_usage_int, - read_usage_value, ) from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.prompt_target.common.target_configuration import TargetConfiguration from pyrit.prompt_target.common.utils import ( build_empty_truncated_response, limit_requests_per_minute, - set_response_metadata, - set_token_usage_metadata, validate_temperature, validate_top_p, - warn_truncated_response, ) -from pyrit.prompt_target.openai.openai_error_handling import _is_content_filter_error +from pyrit.prompt_target.openai._response_adapter import ResponsesResponseAdapter +from pyrit.prompt_target.openai._response_adapter import token_usage_from_responses as token_usage_from_responses from pyrit.prompt_target.openai.openai_target import OpenAITarget if TYPE_CHECKING: @@ -82,47 +77,6 @@ class MessagePieceType(str, Enum): MCP_APPROVAL_REQUEST = "mcp_approval_request" -def token_usage_from_responses(usage: Any) -> TokenUsage: - """ - Build a ``TokenUsage`` from a Responses API ``usage`` payload. - - The Responses API reports usage under different names than Chat Completions -- top-level - ``input_tokens`` / ``output_tokens`` / ``total_tokens`` with ``input_tokens_details`` and - ``output_tokens_details`` breakdowns -- so the field names are resolved here rather than by - ``token_usage_from_chat_completion``. Both parsers share the format-agnostic reads - (``read_usage_value`` / ``read_usage_int``), so a partial usage payload contributes only the - counts the provider actually reports. ``total_tokens`` is derived when the provider omits it. - - Args: - usage (Any): The Responses API usage object. - - Returns: - TokenUsage: The parsed token usage. - """ - input_details = read_usage_value(source=usage, name="input_tokens_details") - output_details = read_usage_value(source=usage, name="output_tokens_details") - - input_tokens = read_usage_int(source=usage, name="input_tokens") - output_tokens = read_usage_int(source=usage, name="output_tokens") - total_tokens = read_usage_int(source=usage, name="total_tokens") - if total_tokens is None and input_tokens is not None and output_tokens is not None: - total_tokens = input_tokens + output_tokens - - extra: dict[str, int] = {} - cache_write_tokens = read_usage_int(source=input_details, name="cache_write_tokens") - if cache_write_tokens is not None: - extra["cache_write_tokens"] = cache_write_tokens - - return TokenUsage( - input_tokens=input_tokens, - output_tokens=output_tokens, - total_tokens=total_tokens, - reasoning_tokens=read_usage_int(source=output_details, name="reasoning_tokens"), - cached_tokens=read_usage_int(source=input_details, name="cached_tokens"), - extra=extra, - ) - - class OpenAIResponseTarget(OpenAITarget): """ Enables communication with endpoints that support the OpenAI Response API. @@ -152,6 +106,7 @@ class OpenAIResponseTarget(OpenAITarget): ), ) ) + _response_adapter = ResponsesResponseAdapter() @forward_init_parameters def __init__( @@ -494,165 +449,6 @@ def _build_text_format(self, json_config: JsonResponseConfig) -> dict[str, Any] logger.info("Using json_object format without schema - consider providing a schema for better results") return {"format": {"type": "json_object"}} - def _check_content_filter(self, response: Any) -> bool: - """ - Check if a Response API response has a content filter error. - - The Responses API signals content filtering in two ways: - 1. Via ``response.error`` with a content_filter code (older/alternative path) - 2. Via ``response.status == "incomplete"`` with - ``response.incomplete_details.reason == "content_filter"`` - - Args: - response: A Response object from the OpenAI SDK. - - Returns: - True if content was filtered, False otherwise. - """ - # Path 1: error-based detection (e.g., error.code == "content_filter") - if hasattr(response, "error") and response.error is not None: - response_dict = response.model_dump() - if _is_content_filter_error(response_dict): - return True - - # Path 2: incomplete status with content_filter reason - if getattr(response, "status", None) == "incomplete": - incomplete_details = getattr(response, "incomplete_details", None) - if incomplete_details and getattr(incomplete_details, "reason", None) == "content_filter": - return True - - return False - - def _extract_partial_content(self, response: Any) -> str | None: - """ - Extract partial content from a Response API response that was content-filtered. - - When the Responses API triggers a content filter, the response may contain partial - output in ``response.output`` message sections with ``status='completed'``. Messages - with ``status='incomplete'`` typically contain refusal text and are excluded. - - Args: - response: A Response object from the OpenAI SDK. - - Returns: - The partial text content from completed output messages, or None if no - partial content was generated. - """ - try: - if not hasattr(response, "output") or not response.output: - return None - parts: list[str] = [] - for section in response.output: - if getattr(section, "type", None) != MessagePieceType.MESSAGE: - continue - # Only include completed messages — incomplete messages contain refusal text - if getattr(section, "status", None) != "completed": - continue - content = getattr(section, "content", None) - parts.extend( - content_item.text - for content_item in content or [] - if isinstance(content_item, ResponseOutputText) and content_item.text - ) - return "\n".join(parts) if parts else None - except (AttributeError, IndexError, TypeError): - return None - - def _capture_response_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: - """ - Record token usage, ``status`` and ``incomplete_reason`` from a Responses API response. - - The Responses API reports why generation stopped as ``status`` plus, when the status is - ``incomplete``, ``incomplete_details.reason`` (for example ``max_output_tokens`` or - ``content_filter``). Together they are this format's equivalent of Chat Completions' - ``finish_reason``. - - Args: - response: A Response object from the OpenAI SDK, or the synthetic stand-in used when - the SDK raises on a content filter. - pieces (list[MessagePiece]): The constructed response pieces. - """ - if not pieces: - return - - usage = getattr(response, "usage", None) - set_token_usage_metadata(pieces=pieces, usage=token_usage_from_responses(usage) if usage is not None else None) - - status = getattr(response, "status", None) - incomplete_details = getattr(response, "incomplete_details", None) - incomplete_reason = getattr(incomplete_details, "reason", None) if incomplete_details else None - set_response_metadata(pieces=pieces, status=status, incomplete_reason=incomplete_reason) - - def _validate_response(self, response: Response, request: MessagePiece) -> None: - """ - Validate a Response API response for errors. - - Checks for: - - Error responses (excluding content filtering which is checked separately) - - Truncation at the token limit (``max_output_tokens``), which is warned about, not raised - - Invalid status - - Empty output - - Truncation is treated as valid, with a warning, so that - ``_construct_message_from_response_async`` can preserve any completed output (reasoning, - partial text) or fall back to a graceful empty response. Genuinely empty responses (no - truncation) are raised so the retry logic can attempt to get a complete response. Content - filter responses are handled separately by ``_check_content_filter``. - - Args: - response: The Response object from the OpenAI SDK. - request: The original request MessagePiece. - - Raises: - PyritException: For unexpected response structures or errors. - EmptyResponseException: When the API returns no valid output (and was not truncated). - """ - # Check for error response - error is a ResponseError object or None - # (content_filter is handled by _check_content_filter) - if response.error is not None and response.error.code != "content_filter": - raise PyritException(message=f"Response error: {response.error.code} - {response.error.message}") - - # Truncation: the model hit max_output_tokens. Mirroring OpenAIChatTarget's handling of - # finish_reason == "length", warn instead of raising so the run continues -- reasoning models - # can spend the whole budget on hidden reasoning before emitting a visible answer, and a low - # limit may be a deliberate configuration. Construction preserves any completed output - # (reasoning, partial text) and falls back to a graceful empty response. - if self._is_truncated_response(response): - warn_truncated_response( - signal="status='incomplete', reason='max_output_tokens'", - limit_parameter="max_output_tokens", - ) - return - - # Check status - should be "completed" for successful responses - if response.status != "completed": - raise PyritException(message=f"Unexpected status: {response.status}") - - # Check for empty output - if not response.output: - logger.error("The response returned no valid output.") - raise EmptyResponseException(message="The response returned an empty response.") - - def _is_truncated_response(self, response: Response) -> bool: - """ - Return True if the response was cut off by the ``max_output_tokens`` limit. - - The Responses API signals truncation via ``status == "incomplete"`` with - ``incomplete_details.reason == "max_output_tokens"`` (``content_filter`` is handled - separately by ``_check_content_filter``). - - Args: - response: A Response object from the OpenAI SDK. - - Returns: - bool: True if the response was truncated at the token limit, False otherwise. - """ - if response.status != "incomplete": - return False - incomplete_details = response.incomplete_details - reason = incomplete_details.reason if incomplete_details else None - return reason == "max_output_tokens" - async def _construct_message_from_response_async(self, response: Response, request: MessagePiece) -> Message: """ Construct a Message from a Response API response. diff --git a/pyrit/prompt_target/openai/openai_target.py b/pyrit/prompt_target/openai/openai_target.py index aea513a065..3f4a292970 100644 --- a/pyrit/prompt_target/openai/openai_target.py +++ b/pyrit/prompt_target/openai/openai_target.py @@ -32,6 +32,7 @@ from pyrit.prompt_target.common.prompt_target import AuthMode, PromptTarget from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.prompt_target.common.target_configuration import TargetConfiguration +from pyrit.prompt_target.openai._response_adapter import NoOpOpenAIResponseAdapter, OpenAIResponseAdapter from pyrit.prompt_target.openai.openai_error_handling import ( _extract_error_payload, _extract_request_id_from_exception, @@ -56,6 +57,7 @@ class OpenAITarget(PromptTarget): _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( capabilities=TargetCapabilities(supports_multi_message_pieces=True) ) + _response_adapter: OpenAIResponseAdapter[Any] = NoOpOpenAIResponseAdapter() # OpenAI-family targets can mint an Entra ID token for a recognized Azure # endpoint (see ``is_azure_openai_endpoint``), so they support both modes. @@ -537,7 +539,7 @@ def _check_content_filter(self, response: Any) -> bool: Returns: bool: True if content filter detected, False otherwise. """ - return False + return self._response_adapter.is_content_filter(response=response) def _handle_content_filter_response(self, response: Any, request: MessagePiece) -> Message: """ @@ -589,7 +591,7 @@ def _extract_partial_content(self, response: Any) -> str | None: Returns: The partial text content, or None if no content was generated. """ - return None + return self._response_adapter.extract_partial_content(response=response) def _capture_response_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: """ @@ -608,6 +610,7 @@ def _capture_response_metadata(self, *, response: Any, pieces: list[MessagePiece no usage or completion data, so implementations must tolerate missing attributes. pieces (list[MessagePiece]): The constructed response pieces. """ + self._response_adapter.capture_metadata(response=response, pieces=pieces) def _validate_response(self, response: Any, request: MessagePiece) -> None: """ @@ -624,6 +627,7 @@ def _validate_response(self, response: Any, request: MessagePiece) -> None: Raises: Various exceptions for validation failures. """ + self._response_adapter.validate(response=response, is_truncated=self._is_truncated_response(response)) def _is_truncated_response(self, response: Any) -> bool: """ @@ -642,7 +646,7 @@ def _is_truncated_response(self, response: Any) -> bool: Returns: bool: True if the response was truncated at the token limit, False otherwise. """ - return False + return self._response_adapter.is_truncated(response=response) @abstractmethod def _set_openai_env_configuration_vars(self) -> None: diff --git a/tests/unit/prompt_target/target/test_openai_response_adapters.py b/tests/unit/prompt_target/target/test_openai_response_adapters.py new file mode 100644 index 0000000000..43136a226b --- /dev/null +++ b/tests/unit/prompt_target/target/test_openai_response_adapters.py @@ -0,0 +1,158 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from unittest.mock import MagicMock, patch + +import pytest +from openai.types.chat import ChatCompletion +from openai.types.chat.chat_completion import Choice +from openai.types.chat.chat_completion_message import ChatCompletionMessage +from openai.types.completion import Completion +from openai.types.completion_choice import CompletionChoice +from openai.types.completion_usage import CompletionUsage +from openai.types.responses import Response, ResponseOutputMessage, ResponseOutputText + +from pyrit.exceptions import PyritException +from pyrit.models import MessagePiece +from pyrit.prompt_target.openai._response_adapter import ( + ChatCompletionsResponseAdapter, + CompletionsResponseAdapter, + ResponsesResponseAdapter, +) +from pyrit.prompt_target.openai.openai_chat_target import OpenAIChatTarget +from pyrit.prompt_target.openai.openai_completion_target import OpenAICompletionTarget +from pyrit.prompt_target.openai.openai_response_target import OpenAIResponseTarget + + +def _piece() -> MessagePiece: + return MessagePiece(role="assistant", original_value="response") + + +def _chat_response(*, content: str | None, finish_reason: str) -> ChatCompletion: + return ChatCompletion( + id="chat-1", + choices=[ + Choice( + finish_reason=finish_reason, + index=0, + message=ChatCompletionMessage(role="assistant", content=content), + ) + ], + created=0, + model="gpt-4o", + object="chat.completion", + usage=CompletionUsage(prompt_tokens=5, completion_tokens=3, total_tokens=8), + ) + + +def _responses_response( + *, + status: str, + incomplete_reason: str | None = None, + text: str = "partial", +) -> MagicMock: + response = MagicMock(spec=Response) + response.error = None + response.status = status + response.incomplete_details = MagicMock(reason=incomplete_reason) if incomplete_reason else None + response.output = [ + ResponseOutputMessage( + id="message-1", + content=[ResponseOutputText(annotations=[], text=text, type="output_text")], + role="assistant", + status="completed", + type="message", + ) + ] + response.usage = MagicMock( + input_tokens=5, + output_tokens=3, + total_tokens=8, + input_tokens_details=None, + output_tokens_details=None, + ) + return response + + +def test_targets_select_explicit_response_format_adapters(): + assert isinstance(OpenAIChatTarget._response_adapter, ChatCompletionsResponseAdapter) + assert isinstance(OpenAICompletionTarget._response_adapter, CompletionsResponseAdapter) + assert isinstance(OpenAIResponseTarget._response_adapter, ResponsesResponseAdapter) + + +def test_chat_completions_adapter_contract(): + adapter = ChatCompletionsResponseAdapter() + filtered = _chat_response(content="partial", finish_reason="content_filter") + truncated = _chat_response(content="", finish_reason="length") + malformed = _chat_response(content="ignored", finish_reason="stop") + malformed.choices = [] + + assert adapter.is_content_filter(response=filtered) is True + assert adapter.extract_partial_content(response=filtered) == "partial" + assert adapter.is_truncated(response=truncated) is True + adapter.validate(response=truncated, is_truncated=adapter.is_truncated(response=truncated)) + with pytest.raises(PyritException, match="No choices returned"): + adapter.validate(response=malformed, is_truncated=adapter.is_truncated(response=malformed)) + + piece = _piece() + adapter.capture_metadata(response=filtered, pieces=[piece]) + assert piece.prompt_metadata["finish_reason"] == "content_filter" + assert piece.prompt_metadata["token_usage_total_tokens"] == 8 + + +def test_completions_adapter_preserves_legacy_contract(): + adapter = CompletionsResponseAdapter() + response = Completion( + id="completion-1", + object="text_completion", + created=0, + model="gpt-3.5-turbo-instruct", + choices=[CompletionChoice(finish_reason="content_filter", index=0, text="partial")], + usage=CompletionUsage(prompt_tokens=5, completion_tokens=3, total_tokens=8), + ) + + assert adapter.is_content_filter(response=response) is False + assert adapter.extract_partial_content(response=response) is None + assert adapter.is_truncated(response=response) is False + adapter.validate(response=response, is_truncated=adapter.is_truncated(response=response)) + + piece = _piece() + adapter.capture_metadata(response=response, pieces=[piece]) + assert piece.prompt_metadata["finish_reason"] == "content_filter" + assert piece.prompt_metadata["token_usage_total_tokens"] == 8 + + +def test_responses_adapter_contract(): + adapter = ResponsesResponseAdapter() + filtered = _responses_response(status="incomplete", incomplete_reason="content_filter") + truncated = _responses_response(status="incomplete", incomplete_reason="max_output_tokens", text="") + malformed = _responses_response(status="failed") + + assert adapter.is_content_filter(response=filtered) is True + assert adapter.extract_partial_content(response=filtered) == "partial" + assert adapter.is_truncated(response=truncated) is True + adapter.validate(response=truncated, is_truncated=adapter.is_truncated(response=truncated)) + with pytest.raises(PyritException, match="Unexpected status: failed"): + adapter.validate(response=malformed, is_truncated=adapter.is_truncated(response=malformed)) + + piece = _piece() + adapter.capture_metadata(response=filtered, pieces=[piece]) + assert piece.prompt_metadata["status"] == "incomplete" + assert piece.prompt_metadata["incomplete_reason"] == "content_filter" + assert piece.prompt_metadata["token_usage_total_tokens"] == 8 + + +def test_chat_target_validation_honors_truncation_override(): + target = object.__new__(OpenAIChatTarget) + response = _chat_response(content="", finish_reason="stop") + + with patch.object(target, "_is_truncated_response", return_value=True): + target._validate_response(response, _piece()) + + +def test_responses_target_validation_honors_truncation_override(): + target = object.__new__(OpenAIResponseTarget) + response = _responses_response(status="failed") + + with patch.object(target, "_is_truncated_response", return_value=True): + target._validate_response(response, _piece()) From e5521d98f50d7964130af16ca4aad9117c697533 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Thu, 13 Aug 2026 14:22:13 -0700 Subject: [PATCH 2/2] MAINT: Refine OpenAI response adapters Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 7b3a8d42-bfe2-4b5e-90ae-fc211028c2ad --- .../prompt_target/openai/_response_adapter.py | 56 +++++++------------ pyrit/prompt_target/openai/openai_target.py | 6 +- .../target/test_openai_response_adapters.py | 20 ++++++- 3 files changed, 40 insertions(+), 42 deletions(-) diff --git a/pyrit/prompt_target/openai/_response_adapter.py b/pyrit/prompt_target/openai/_response_adapter.py index 4b3be7c28f..8b8d629fc8 100644 --- a/pyrit/prompt_target/openai/_response_adapter.py +++ b/pyrit/prompt_target/openai/_response_adapter.py @@ -2,7 +2,7 @@ # Licensed under the MIT license. import logging -from typing import Any, Protocol, TypeVar +from typing import Any, Generic, TypeVar from openai.types.chat import ChatCompletion from openai.types.responses import Response, ResponseOutputText @@ -29,56 +29,40 @@ ResponseT = TypeVar("ResponseT", contravariant=True) -class OpenAIResponseAdapter(Protocol[ResponseT]): - """The response-format contract used by ``OpenAITarget``.""" +class OpenAIResponseAdapter(Generic[ResponseT]): + """Base response-format behavior used by ``OpenAITarget``.""" - def is_content_filter(self, *, response: ResponseT) -> bool: + def is_content_filtered(self, *, response: ResponseT) -> bool: """Return whether the response was blocked by a content filter.""" - ... + return False def extract_partial_content(self, *, response: ResponseT) -> str | None: - """Extract content emitted before a content filter stopped generation.""" - ... + """ + Extract content emitted before a content filter stopped generation. + + Args: + response (ResponseT): The provider response. + + Returns: + str | None: Partial content, if available. + """ + return None def capture_metadata(self, *, response: ResponseT, pieces: list[MessagePiece]) -> None: """Copy provider response metadata to PyRIT response pieces.""" - ... def validate(self, *, response: ResponseT, is_truncated: bool) -> None: """Validate the provider response.""" - ... def is_truncated(self, *, response: ResponseT) -> bool: """Return whether generation stopped at the output-token limit.""" - ... - - -class NoOpOpenAIResponseAdapter: - """Default behavior for OpenAI targets without a structured response format.""" - - def is_content_filter(self, *, response: Any) -> bool: - """Return False because this format has no content-filter signal.""" - return False - - def extract_partial_content(self, *, response: Any) -> str | None: - """Return no partial content because this format has no extraction rule.""" - return None - - def capture_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: - """Leave response metadata unchanged.""" - - def validate(self, *, response: Any, is_truncated: bool) -> None: - """Accept the response without format-specific validation.""" - - def is_truncated(self, *, response: Any) -> bool: - """Return False because this format has no truncation signal.""" return False -class ChatCompletionsResponseAdapter: +class ChatCompletionsResponseAdapter(OpenAIResponseAdapter[ChatCompletion]): """Response behavior for the OpenAI Chat Completions wire format.""" - def is_content_filter(self, *, response: ChatCompletion) -> bool: + def is_content_filtered(self, *, response: ChatCompletion) -> bool: """Return whether ``finish_reason`` reports a content filter.""" return is_content_filter_response(response) @@ -110,7 +94,7 @@ def is_truncated(self, *, response: ChatCompletion) -> bool: return get_finish_reason(response=response) == "length" -class CompletionsResponseAdapter(NoOpOpenAIResponseAdapter): +class CompletionsResponseAdapter(OpenAIResponseAdapter[Any]): """Response behavior for the legacy OpenAI Completions wire format.""" def capture_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None: @@ -123,10 +107,10 @@ def capture_metadata(self, *, response: Any, pieces: list[MessagePiece]) -> None set_response_metadata(pieces=[piece], finish_reason=getattr(choice, "finish_reason", None)) -class ResponsesResponseAdapter: +class ResponsesResponseAdapter(OpenAIResponseAdapter[Response]): """Response behavior for the OpenAI Responses API wire format.""" - def is_content_filter(self, *, response: Response) -> bool: + def is_content_filtered(self, *, response: Response) -> bool: """Return whether the response reports content filtering.""" error = getattr(response, "error", None) if error is not None and _is_content_filter_error(response.model_dump()): diff --git a/pyrit/prompt_target/openai/openai_target.py b/pyrit/prompt_target/openai/openai_target.py index 3f4a292970..6e6efbebaa 100644 --- a/pyrit/prompt_target/openai/openai_target.py +++ b/pyrit/prompt_target/openai/openai_target.py @@ -32,7 +32,7 @@ from pyrit.prompt_target.common.prompt_target import AuthMode, PromptTarget from pyrit.prompt_target.common.target_capabilities import TargetCapabilities from pyrit.prompt_target.common.target_configuration import TargetConfiguration -from pyrit.prompt_target.openai._response_adapter import NoOpOpenAIResponseAdapter, OpenAIResponseAdapter +from pyrit.prompt_target.openai._response_adapter import OpenAIResponseAdapter from pyrit.prompt_target.openai.openai_error_handling import ( _extract_error_payload, _extract_request_id_from_exception, @@ -57,7 +57,7 @@ class OpenAITarget(PromptTarget): _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( capabilities=TargetCapabilities(supports_multi_message_pieces=True) ) - _response_adapter: OpenAIResponseAdapter[Any] = NoOpOpenAIResponseAdapter() + _response_adapter: OpenAIResponseAdapter[Any] = OpenAIResponseAdapter() # OpenAI-family targets can mint an Entra ID token for a recognized Azure # endpoint (see ``is_azure_openai_endpoint``), so they support both modes. @@ -539,7 +539,7 @@ def _check_content_filter(self, response: Any) -> bool: Returns: bool: True if content filter detected, False otherwise. """ - return self._response_adapter.is_content_filter(response=response) + return self._response_adapter.is_content_filtered(response=response) def _handle_content_filter_response(self, response: Any, request: MessagePiece) -> Message: """ diff --git a/tests/unit/prompt_target/target/test_openai_response_adapters.py b/tests/unit/prompt_target/target/test_openai_response_adapters.py index 43136a226b..91c53e9807 100644 --- a/tests/unit/prompt_target/target/test_openai_response_adapters.py +++ b/tests/unit/prompt_target/target/test_openai_response_adapters.py @@ -17,6 +17,7 @@ from pyrit.prompt_target.openai._response_adapter import ( ChatCompletionsResponseAdapter, CompletionsResponseAdapter, + OpenAIResponseAdapter, ResponsesResponseAdapter, ) from pyrit.prompt_target.openai.openai_chat_target import OpenAIChatTarget @@ -80,6 +81,19 @@ def test_targets_select_explicit_response_format_adapters(): assert isinstance(OpenAIResponseTarget._response_adapter, ResponsesResponseAdapter) +def test_base_adapter_uses_no_op_defaults(): + adapter = OpenAIResponseAdapter[object]() + response = object() + piece = _piece() + + assert adapter.is_content_filtered(response=response) is False + assert adapter.extract_partial_content(response=response) is None + assert adapter.is_truncated(response=response) is False + adapter.capture_metadata(response=response, pieces=[piece]) + adapter.validate(response=response, is_truncated=False) + assert piece.prompt_metadata == {} + + def test_chat_completions_adapter_contract(): adapter = ChatCompletionsResponseAdapter() filtered = _chat_response(content="partial", finish_reason="content_filter") @@ -87,7 +101,7 @@ def test_chat_completions_adapter_contract(): malformed = _chat_response(content="ignored", finish_reason="stop") malformed.choices = [] - assert adapter.is_content_filter(response=filtered) is True + assert adapter.is_content_filtered(response=filtered) is True assert adapter.extract_partial_content(response=filtered) == "partial" assert adapter.is_truncated(response=truncated) is True adapter.validate(response=truncated, is_truncated=adapter.is_truncated(response=truncated)) @@ -111,7 +125,7 @@ def test_completions_adapter_preserves_legacy_contract(): usage=CompletionUsage(prompt_tokens=5, completion_tokens=3, total_tokens=8), ) - assert adapter.is_content_filter(response=response) is False + assert adapter.is_content_filtered(response=response) is False assert adapter.extract_partial_content(response=response) is None assert adapter.is_truncated(response=response) is False adapter.validate(response=response, is_truncated=adapter.is_truncated(response=response)) @@ -128,7 +142,7 @@ def test_responses_adapter_contract(): truncated = _responses_response(status="incomplete", incomplete_reason="max_output_tokens", text="") malformed = _responses_response(status="failed") - assert adapter.is_content_filter(response=filtered) is True + assert adapter.is_content_filtered(response=filtered) is True assert adapter.extract_partial_content(response=filtered) == "partial" assert adapter.is_truncated(response=truncated) is True adapter.validate(response=truncated, is_truncated=adapter.is_truncated(response=truncated))