From 92fe5c7f873eb7505e40c6846892ecca175acfe1 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Thu, 1 Oct 2026 21:53:00 +0000 Subject: [PATCH 1/3] FIX: Bound backend request sizes and field lengths The backend accepted request bodies, strings, lists, and maps of any size, silently dropped label filters without a key:value separator, and returned 500 for turn filters too large for the database. Request bodies over 100 MiB now receive 413 with the same problem response whether the size comes from Content-Length or is reached while the body streams, and URLs over 8 KiB receive 414. Identifiers, names, labels, filters, cursors, configuration files, and initializer scripts have length or item limits shared from the backend common models, and the scenario run requests get the same limits inline (technique tokens allow converter modifiers). Values over a limit, label filters without a separator or with a key or value over the label limits, negative cutoff indexes, and turn filters above 10,000 receive 422. Prompt content (message pieces, system prompts, and converter preview input) and free-form values (prompt metadata and parameter values) are limited only by the body size, so long prompts and base64 media keep working. --- pyrit/backend/README.md | 10 ++ pyrit/backend/main.py | 4 + pyrit/backend/middleware/error_handlers.py | 14 +++ pyrit/backend/middleware/request_size.py | 88 +++++++++++++ pyrit/backend/models/attacks.py | 54 +++++--- pyrit/backend/models/common.py | 54 +++++++- pyrit/backend/models/configuration.py | 12 +- pyrit/backend/models/converters.py | 9 +- pyrit/backend/models/initializers.py | 8 +- pyrit/backend/models/scores.py | 4 +- pyrit/backend/models/targets.py | 8 +- pyrit/backend/routes/attacks.py | 47 ++++--- pyrit/backend/routes/common.py | 11 +- pyrit/backend/routes/configuration.py | 6 +- pyrit/backend/routes/converters.py | 6 +- pyrit/backend/routes/initializers.py | 8 +- pyrit/backend/routes/labels.py | 7 +- pyrit/backend/routes/scenarios.py | 29 +++-- pyrit/backend/routes/targets.py | 6 +- pyrit/models/catalog/scenario.py | 64 +++++++--- tests/unit/backend/test_api_routes.py | 80 +++++++++--- tests/unit/backend/test_common_models.py | 137 +++++++++++++++++++++ tests/unit/backend/test_request_size.py | 93 ++++++++++++++ tests/unit/models/test_scenario_request.py | 61 ++++++++- 24 files changed, 697 insertions(+), 123 deletions(-) create mode 100644 pyrit/backend/middleware/request_size.py create mode 100644 tests/unit/backend/test_request_size.py diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index c9e0c3420d..f2e6ca856c 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -61,6 +61,16 @@ Concurrent sends or `send=false` appends to the same conversation receive **409* exceeding the admission limit receives **429**. Neither response starts a send or appends a message. No background submission or status API is introduced. +### Request Limits + +The backend reads at most 100 MiB of a request body; larger bodies receive **413**, and +API requests with URLs over 8 KiB receive **414**. Identifiers, names, labels, filters, +cursors, configuration files, and initializer scripts also have length or item limits, +listed in the OpenAPI schema; values over a limit receive **422**. Prompt content (message +pieces, system prompts, and converter preview input) and free-form values (prompt +metadata, scenario and initializer arguments, and target and converter parameter values) +are limited only by the body size, so long prompts and base64 media keep working. + ## Strict Lockstep Compatibility The backend, CLI, and frontend bundle use one stamped identity: diff --git a/pyrit/backend/main.py b/pyrit/backend/main.py index 0be8af50ca..9413195912 100644 --- a/pyrit/backend/main.py +++ b/pyrit/backend/main.py @@ -24,6 +24,7 @@ from pyrit.backend.middleware import RequestIdMiddleware, SecurityHeadersMiddleware, register_error_handlers from pyrit.backend.middleware.auth import EntraAuthMiddleware from pyrit.backend.middleware.compatibility import CompatibilityAPI, CompatibilityMiddleware +from pyrit.backend.middleware.request_size import RequestSizeLimitMiddleware from pyrit.backend.middleware.runtime import RuntimeAdmissionMiddleware from pyrit.backend.routes import ( attacks, @@ -90,6 +91,9 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: # Register RFC 7807 error handlers register_error_handlers(app) +# Innermost so route handlers read the body directly through its size check; a BaseHTTPMiddleware +# between them would turn the 413 raised while streaming into a body-parsing error. +app.add_middleware(RequestSizeLimitMiddleware) app.add_middleware(RuntimeAdmissionMiddleware) # Microsoft Graph-backed authentication (PKCE — no client secrets needed) diff --git a/pyrit/backend/middleware/error_handlers.py b/pyrit/backend/middleware/error_handlers.py index 44f4d340ed..a7ad9a7429 100644 --- a/pyrit/backend/middleware/error_handlers.py +++ b/pyrit/backend/middleware/error_handlers.py @@ -11,6 +11,7 @@ from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse +from pyrit.backend.middleware.request_size import RequestTooLargeError, request_too_large_response from pyrit.backend.models.common import FieldError, ProblemDetail logger = logging.getLogger(__name__) @@ -55,6 +56,19 @@ async def validation_exception_handler( # pyrit-async-suffix-exempt content=problem.model_dump(exclude_none=True), ) + @app.exception_handler(RequestTooLargeError) + async def request_too_large_handler( # pyrit-async-suffix-exempt + request: Request, + exc: RequestTooLargeError, + ) -> JSONResponse: + """ + Handle a request body that grew past the size limit while it was read. + + Returns: + JSONResponse: The same RFC 7807 problem response the size-limit middleware sends. + """ + return request_too_large_response(status=status.HTTP_413_CONTENT_TOO_LARGE, title="Content Too Large") + @app.exception_handler(ValueError) async def value_error_handler( # pyrit-async-suffix-exempt request: Request, diff --git a/pyrit/backend/middleware/request_size.py b/pyrit/backend/middleware/request_size.py new file mode 100644 index 0000000000..28cfccce03 --- /dev/null +++ b/pyrit/backend/middleware/request_size.py @@ -0,0 +1,88 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Reject requests whose URL or body is larger than the backend accepts.""" + +from starlette.datastructures import Headers +from starlette.exceptions import HTTPException +from starlette.responses import JSONResponse +from starlette.types import ASGIApp, Message, Receive, Scope, Send + + +class RequestTooLargeError(HTTPException): + """Raised when a request body grows past the size limit while it is read.""" + + def __init__(self) -> None: + """Create the 413 error.""" + super().__init__(status_code=413, detail="The request is larger than the backend accepts.") + + +def request_too_large_response(*, status: int, title: str) -> JSONResponse: + """ + Build the RFC 7807 problem response for a request over a size limit. + + Returns: + JSONResponse: The problem response. + """ + return JSONResponse( + status_code=status, + media_type="application/problem+json", + content={ + "type": "/errors/request-too-large", + "title": title, + "status": status, + "detail": "The request is larger than the backend accepts.", + }, + ) + + +class RequestSizeLimitMiddleware: + """Return 414 for oversized URLs and 413 for oversized bodies before route handlers use them.""" + + # Large enough for base64-encoded media attachments. + MAX_BODY_BYTES: int = 100 * 1024 * 1024 + MAX_URL_LENGTH: int = 8 * 1024 + + def __init__(self, app: ASGIApp) -> None: + """Wrap the downstream application.""" + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + """Reject oversized requests and cap the body bytes the application can read.""" + if scope["type"] != "http": + await self.app(scope, receive, send) + return + path = scope.get("raw_path") or scope["path"].encode() + if len(path) + len(scope.get("query_string", b"")) > self.MAX_URL_LENGTH: + await self._reject_async(scope=scope, receive=receive, send=send, status=414, title="URI Too Long") + return + content_length = Headers(scope=scope).get("content-length", "") + if content_length.isdigit() and int(content_length) > self.MAX_BODY_BYTES: + await self._reject_async(scope=scope, receive=receive, send=send, status=413, title="Content Too Large") + return + await self.app(scope, self._limit_body(receive), send) + + def _limit_body(self, receive: Receive) -> Receive: + """ + Wrap ``receive`` so bodies sent without a reliable ``Content-Length`` are also capped. + + Returns: + Receive: A receive callable that raises a 413 error once the cap is exceeded. + """ + received = 0 + + async def limited_receive_async() -> Message: + nonlocal received + message = await receive() + if message["type"] == "http.request": + received += len(message.get("body", b"")) + if received > self.MAX_BODY_BYTES: + raise RequestTooLargeError + return message + + return limited_receive_async + + @staticmethod + async def _reject_async(*, scope: Scope, receive: Receive, send: Send, status: int, title: str) -> None: + """Send an RFC 7807 problem response without calling the application.""" + await request_too_large_response(status=status, title=title)(scope, receive, send) diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index 0cd5a99869..d298c64fc9 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -15,7 +15,13 @@ from pydantic import BaseModel, Field, computed_field, field_serializer, model_validator from pyrit.backend.models._media import build_filename, infer_mime_type -from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.models.common import ( + MAX_ITEMS, + IdentifierStr, + LabelDict, + PaginationInfo, + TextStr, +) from pyrit.models import ( AttackResult, ChatMessageRole, @@ -366,17 +372,19 @@ class MessagePieceRequest(BaseModel): None, description="Final converted value's data type. Defaults to data_type; requires converted_value.", ) - applied_converter_ids: list[str] | None = Field( + applied_converter_ids: list[IdentifierStr] | None = Field( None, + max_length=MAX_ITEMS, description="Registry IDs of converters already applied, in execution order, including duplicates. " "Requires converted_value. Use an empty list for manual edits.", ) - mime_type: str | None = Field(None, description="MIME type for media content") - prompt_metadata: dict[str, Any] | None = Field( + mime_type: IdentifierStr | None = Field(None, description="MIME type for media content") + prompt_metadata: dict[IdentifierStr, Any] | None = Field( None, + max_length=MAX_ITEMS, description="Metadata to attach to the piece (e.g., {'video_id': '...'} for remix mode).", ) - original_prompt_id: str | None = Field( + original_prompt_id: IdentifierStr | None = Field( None, description="ID of the source piece when prepending from an existing conversation. " "Preserves lineage so the new piece traces back to the original.", @@ -412,7 +420,7 @@ class _AttackAttributionInput(BaseModel): operator: str | None = Field(None, max_length=128, description="Operator responsible for the attack") operation: str | None = Field(None, max_length=128, description="Operation associated with the attack") - labels: dict[str, str] | None = Field(None, description="Arbitrary user-defined labels for filtering") + labels: LabelDict | None = Field(None, description="Arbitrary user-defined labels for filtering") @model_validator(mode="before") @classmethod @@ -458,12 +466,14 @@ class CreateAttackRequest(_AttackAttributionInput): supplied in ``labels`` (typically the current operator's labels). """ - name: str | None = Field(None, description="Attack name/label") - target_registry_name: str = Field(..., description="Target registry name to attack") - source_conversation_id: str | None = Field( + name: TextStr | None = Field(None, description="Attack name/label") + target_registry_name: IdentifierStr = Field(..., description="Target registry name to attack") + source_conversation_id: IdentifierStr | None = Field( None, description="Conversation to branch from (clone messages into the new attack)" ) - cutoff_index: int | None = Field(None, description="Include messages up to and including this turn index (0-based)") + cutoff_index: int | None = Field( + None, ge=0, description="Include messages up to and including this turn index (0-based)" + ) system_prompt: str | None = Field( None, description="System prompt lowered to a single system-role message at the front of the conversation. " @@ -494,7 +504,7 @@ class UpdateAttackRequest(BaseModel): default=None, description="Updated attack outcome", ) - objective: str | None = Field(default=None, description="Shared objective for all conversations in the attack") + objective: TextStr | None = Field(default=None, description="Shared objective for all conversations in the attack") @model_validator(mode="after") def _validate_update(self) -> "UpdateAttackRequest": @@ -546,8 +556,10 @@ class CreateConversationRequest(BaseModel): the cutoff turn, preserving tracking relationships (original_prompt_id). """ - source_conversation_id: str | None = Field(None, description="Conversation to branch from") - cutoff_index: int | None = Field(None, description="Include messages up to and including this turn index (0-based)") + source_conversation_id: IdentifierStr | None = Field(None, description="Conversation to branch from") + cutoff_index: int | None = Field( + None, ge=0, description="Include messages up to and including this turn index (0-based)" + ) class CreateConversationResponse(BaseModel): @@ -560,7 +572,7 @@ class CreateConversationResponse(BaseModel): class UpdateMainConversationRequest(BaseModel): """Request to update the main conversation of an attack result.""" - conversation_id: str = Field(..., description="The conversation to promote to main") + conversation_id: IdentifierStr = Field(..., description="The conversation to promote to main") class UpdateMainConversationResponse(BaseModel): @@ -579,19 +591,22 @@ class UpdateMainConversationResponse(BaseModel): class ConverterConfigurationRequest(BaseModel): """Registry-backed converter configuration for one ordered pipeline.""" - converter_ids: list[str] = Field( + converter_ids: list[IdentifierStr] = Field( ..., min_length=1, + max_length=MAX_ITEMS, description="Converter instance IDs to apply in order.", ) indexes_to_apply: list[Annotated[int, Field(ge=0)]] | None = Field( None, min_length=1, + max_length=MAX_ITEMS, description="Zero-based message piece indexes to which this pipeline applies. Defaults to all indexes.", ) prompt_data_types_to_apply: list[PromptDataType] | None = Field( None, min_length=1, + max_length=MAX_ITEMS, description="Prompt data types to which this pipeline applies. Defaults to all data types.", ) @@ -611,25 +626,28 @@ class AddMessageRequest(BaseModel): default=True, description="If True, send to target and wait for response. If False, just store in memory.", ) - target_registry_name: str | None = Field( + target_registry_name: IdentifierStr | None = Field( None, description="Target registry name. Required when send=True so the backend knows which target to use.", ) - converter_ids: list[str] | None = Field( + converter_ids: list[IdentifierStr] | None = Field( None, + max_length=MAX_ITEMS, description="Deprecated global request converter pipeline. Use request_converter_configurations instead.", ) request_converter_configurations: list[ConverterConfigurationRequest] | None = Field( None, min_length=1, + max_length=MAX_ITEMS, description="Ordered registry-backed converter pipelines to apply to the request.", ) response_converter_configurations: list[ConverterConfigurationRequest] | None = Field( None, min_length=1, + max_length=MAX_ITEMS, description="Ordered registry-backed converter pipelines to apply to the response.", ) - target_conversation_id: str = Field( + target_conversation_id: IdentifierStr = Field( ..., description="The conversation_id to store and send messages under. " "Usually the attack's main conversation, but can be a related conversation.", diff --git a/pyrit/backend/models/common.py b/pyrit/backend/models/common.py index 33e751dbc7..2a18939f92 100644 --- a/pyrit/backend/models/common.py +++ b/pyrit/backend/models/common.py @@ -2,17 +2,63 @@ # Licensed under the MIT license. """ -Common response models for the PyRIT API. +Common models for the PyRIT API. -Includes pagination, error handling (RFC 7807), and shared base models. +Includes pagination, error handling (RFC 7807), shared base models, and request limits. """ -from typing import Any +from typing import Annotated, Any -from pydantic import BaseModel, Field +from pydantic import AfterValidator, BaseModel, Field REGISTRY_INSTANCE_NAME_PATTERN = r"^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$" +# Request limits. Prompt content (message pieces, system prompts, preview input) and free-form +# values (metadata and parameter values) are only limited by the request body size, so long +# prompts and base64 media keep working. +MAX_IDENTIFIER_LENGTH = 256 +MAX_CURSOR_LENGTH = 1_024 +MAX_TEXT_LENGTH = 100_000 +MAX_ITEMS = 100 +MAX_LABEL_KEY_LENGTH = 128 +MAX_LABEL_VALUE_LENGTH = 1_024 +MAX_FILE_CONTENT_LENGTH = 1_048_576 + +IdentifierStr = Annotated[str, Field(max_length=MAX_IDENTIFIER_LENGTH)] +CursorStr = Annotated[str, Field(max_length=MAX_CURSOR_LENGTH)] +TextStr = Annotated[str, Field(max_length=MAX_TEXT_LENGTH)] +LabelDict = Annotated[ + dict[ + Annotated[str, Field(max_length=MAX_LABEL_KEY_LENGTH)], + Annotated[str, Field(max_length=MAX_LABEL_VALUE_LENGTH)], + ], + Field(max_length=MAX_ITEMS), +] + + +def validate_label_filter(value: str) -> str: + """ + Check that a label filter is ``key:value`` within the label key and value limits. + + Returns: + str: The unchanged filter. + + Raises: + ValueError: If the filter has no ``:`` separator or a part is too long. + """ + key, separator, label_value = value.partition(":") + if not separator: + raise ValueError("Label filters must use the key:value format.") + if len(key.strip()) > MAX_LABEL_KEY_LENGTH or len(label_value.strip()) > MAX_LABEL_VALUE_LENGTH: + raise ValueError( + f"Label filter keys are limited to {MAX_LABEL_KEY_LENGTH} characters " + f"and values to {MAX_LABEL_VALUE_LENGTH} characters." + ) + return value + + +LabelFilterStr = Annotated[str, AfterValidator(validate_label_filter)] + class PaginationInfo(BaseModel): """Pagination metadata for list responses.""" diff --git a/pyrit/backend/models/configuration.py b/pyrit/backend/models/configuration.py index 40d724e030..a40d5218fc 100644 --- a/pyrit/backend/models/configuration.py +++ b/pyrit/backend/models/configuration.py @@ -5,6 +5,8 @@ from pydantic import BaseModel, Field +from pyrit.backend.models.common import MAX_FILE_CONTENT_LENGTH, IdentifierStr + class ConfigurationFileContent(BaseModel): """Raw UTF-8 contents of the backend configuration file.""" @@ -21,8 +23,8 @@ class ConfigurationFileContent(BaseModel): class UpdateConfigurationFileRequest(BaseModel): """Replacement contents for the backend configuration file.""" - content: str = Field(..., description="Raw YAML configuration file contents") - version: str = Field(..., description="Version token returned by the latest content read") + content: str = Field(..., max_length=MAX_FILE_CONTENT_LENGTH, description="Raw YAML configuration file contents") + version: IdentifierStr = Field(..., description="Version token returned by the latest content read") class EnvironmentFileContent(BaseModel): @@ -47,11 +49,11 @@ class EnvironmentFileListResponse(BaseModel): class UpdateEnvironmentFileRequest(BaseModel): """Replacement contents for an environment file.""" - content: str = Field(..., description="Raw dotenv file contents") - version: str = Field(..., description="Version token returned by the latest content read") + content: str = Field(..., max_length=MAX_FILE_CONTENT_LENGTH, description="Raw dotenv file contents") + version: IdentifierStr = Field(..., description="Version token returned by the latest content read") class ReinitializeRequest(BaseModel): """Apply saved sources to an idle single-process runtime.""" - version: str + version: IdentifierStr diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 316cedb42b..57fdf4c82e 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -11,7 +11,7 @@ from pydantic import BaseModel, Field -from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN +from pyrit.backend.models.common import MAX_ITEMS, REGISTRY_INSTANCE_NAME_PATTERN, IdentifierStr from pyrit.models import ConverterIdentifier, Parameter, PromptDataType __all__ = [ @@ -89,9 +89,10 @@ class CreateConverterRequest(BaseModel): pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name for the converter instance", ) - type: str = Field(..., description="Converter type (e.g., 'Base64Converter')") - params: dict[str, Any] = Field( + type: IdentifierStr = Field(..., description="Converter type (e.g., 'Base64Converter')") + params: dict[IdentifierStr, Any] = Field( default_factory=dict, + max_length=MAX_ITEMS, description="Converter constructor parameters", ) @@ -117,7 +118,7 @@ class ConverterPreviewRequest(BaseModel): original_value: str = Field(..., description="Text to convert") original_value_data_type: PromptDataType = Field(default="text", description="Data type of original value") - converter_ids: list[str] = Field(..., description="Converter instance IDs to apply") + converter_ids: list[IdentifierStr] = Field(..., max_length=MAX_ITEMS, description="Converter instance IDs to apply") class ConverterPreviewResponse(BaseModel): diff --git a/pyrit/backend/models/initializers.py b/pyrit/backend/models/initializers.py index 4a51fb2dd9..3f37bbfc28 100644 --- a/pyrit/backend/models/initializers.py +++ b/pyrit/backend/models/initializers.py @@ -13,7 +13,7 @@ from pydantic import BaseModel, Field -from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.models.common import MAX_FILE_CONTENT_LENGTH, PaginationInfo from pyrit.models import REGISTRY_NAME_PATTERN from pyrit.models.catalog.initializer import RegisteredInitializer @@ -42,7 +42,11 @@ class RegisterInitializerRequest(BaseModel): pattern=REGISTRY_NAME_PATTERN, description="Registry name for the initializer (e.g., 'my_custom')", ) - script_content: str = Field(..., description="Python source code containing a PyRITInitializer subclass") + script_content: str = Field( + ..., + max_length=MAX_FILE_CONTENT_LENGTH, + description="Python source code containing a PyRITInitializer subclass", + ) class CustomInitializerResponse(BaseModel): diff --git a/pyrit/backend/models/scores.py b/pyrit/backend/models/scores.py index 57c4c0162f..22f4a73873 100644 --- a/pyrit/backend/models/scores.py +++ b/pyrit/backend/models/scores.py @@ -7,6 +7,8 @@ from pydantic import BaseModel, ConfigDict, Field, StrictBool +from pyrit.backend.models.common import TextStr + class ManualScoreRequest(BaseModel): """Request to attach a user-supplied score to a message piece.""" @@ -16,7 +18,7 @@ class ManualScoreRequest(BaseModel): attack_result_id: uuid.UUID = Field(..., description="ID of the attack containing the message") message_id: uuid.UUID = Field(..., description="ID of the message piece to score") value: StrictBool = Field(..., description="Whether the attack objective was achieved") - rationale: str = Field(default="", description="Optional explanation for the score") + rationale: TextStr = Field(default="", description="Optional explanation for the score") update_attack: bool = Field( default=False, description="Whether to make this the attack's human score and update its outcome", diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index ba9759e767..10b87ee9e0 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -12,7 +12,7 @@ from pydantic import BaseModel, Field -from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN, PaginationInfo +from pyrit.backend.models.common import MAX_ITEMS, REGISTRY_INSTANCE_NAME_PATTERN, IdentifierStr, PaginationInfo from pyrit.models import JSONValue, Parameter from pyrit.models.catalog.target import TargetInstance @@ -67,8 +67,10 @@ class CreateTargetRequest(BaseModel): pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name; omitted only for legacy UI compatibility", ) - type: str = Field(..., description="Target type (e.g., 'OpenAIChatTarget')") - params: dict[str, JSONValue] = Field(default_factory=dict, description="Target constructor parameters") + type: IdentifierStr = Field(..., description="Target type (e.g., 'OpenAIChatTarget')") + params: dict[IdentifierStr, JSONValue] = Field( + default_factory=dict, max_length=MAX_ITEMS, description="Target constructor parameters" + ) auth_mode: Literal["api_key", "identity"] = Field( "api_key", description=( diff --git a/pyrit/backend/routes/attacks.py b/pyrit/backend/routes/attacks.py index f5935f4e43..7f65c8c46a 100644 --- a/pyrit/backend/routes/attacks.py +++ b/pyrit/backend/routes/attacks.py @@ -31,7 +31,13 @@ UpdateMainConversationRequest, UpdateMainConversationResponse, ) -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import ( + MAX_ITEMS, + CursorStr, + IdentifierStr, + LabelFilterStr, + ProblemDetail, +) from pyrit.backend.routes.common import parse_label_query_params from pyrit.backend.services.attack_service import AttackObjectiveConflictError, get_attack_service from pyrit.backend.services.manual_send_scheduler import ManualSendConflictError, ManualSendQueueFullError @@ -47,14 +53,16 @@ response_model=AttackListResponse, ) async def list_attacks( # pyrit-async-suffix-exempt - attack_types: list[str] | None = Query( + attack_types: list[IdentifierStr] | None = Query( None, + max_length=MAX_ITEMS, description="Filter by attack type names. May be specified multiple times to OR-match " "across types (e.g. ?attack_types=A&attack_types=B). Case-insensitive. " "Omit to return all attacks regardless of type.", ), - converter_types: list[str] | None = Query( + converter_types: list[IdentifierStr] | None = Query( None, + max_length=MAX_ITEMS, description="Filter by converter type names. May be specified multiple times; " "combination semantics are controlled by converter_types_match " "(e.g. ?converter_types=A&converter_types=B). " @@ -79,21 +87,22 @@ async def list_attacks( # pyrit-async-suffix-exempt None, description="Filter by outcome" ), operator: list[Annotated[str, Field(max_length=128)]] | None = Query( - None, description="Filter by dedicated operator values" + None, max_length=MAX_ITEMS, description="Filter by dedicated operator values" ), operation: list[Annotated[str, Field(max_length=128)]] | None = Query( - None, description="Filter by dedicated operation values" + None, max_length=MAX_ITEMS, description="Filter by dedicated operation values" ), - label: list[str] | None = Query( + label: list[LabelFilterStr] | None = Query( None, + max_length=MAX_ITEMS, description="Filter by labels (format: key:value). May be specified multiple times; " "OR-matched within a key, AND-matched across keys " "(e.g. ?label=op:red&label=op:blue matches op=red OR op=blue).", ), - min_turns: int | None = Query(None, ge=0, description="Filter by minimum executed turns"), - max_turns: int | None = Query(None, ge=0, description="Filter by maximum executed turns"), + min_turns: int | None = Query(None, ge=0, le=10_000, description="Filter by minimum executed turns"), + max_turns: int | None = Query(None, ge=0, le=10_000, description="Filter by maximum executed turns"), limit: int = Query(20, ge=1, le=100, description="Maximum items per page"), - cursor: str | None = Query( + cursor: CursorStr | None = Query( None, description="Opaque pagination cursor returned as next_cursor by the previous page. " "Treat it as opaque and pass it back unmodified. " @@ -232,7 +241,7 @@ async def create_attack(request: CreateAttackRequest) -> CreateAttackResponse: 404: {"model": ProblemDetail, "description": "Attack not found"}, }, ) -async def get_attack(attack_result_id: str) -> AttackSummary: # pyrit-async-suffix-exempt +async def get_attack(attack_result_id: IdentifierStr) -> AttackSummary: # pyrit-async-suffix-exempt """ Get attack details. @@ -262,7 +271,7 @@ async def get_attack(attack_result_id: str) -> AttackSummary: # pyrit-async-suf }, ) async def update_attack( # pyrit-async-suffix-exempt - attack_result_id: str, + attack_result_id: IdentifierStr, request: UpdateAttackRequest, ) -> AttackSummary: """ @@ -293,7 +302,7 @@ async def update_attack( # pyrit-async-suffix-exempt 404: {"model": ProblemDetail, "description": "Attack not found"}, }, ) -async def remove_human_score(attack_result_id: str) -> AttackSummary: # pyrit-async-suffix-exempt +async def remove_human_score(attack_result_id: IdentifierStr) -> AttackSummary: # pyrit-async-suffix-exempt """ Remove the attack's human-score override. @@ -319,8 +328,8 @@ async def remove_human_score(attack_result_id: str) -> AttackSummary: # pyrit-a }, ) async def get_conversation_messages( # pyrit-async-suffix-exempt - attack_result_id: str, - conversation_id: str = Query(..., description="The conversation_id whose messages to return"), + attack_result_id: IdentifierStr, + conversation_id: IdentifierStr = Query(..., description="The conversation_id whose messages to return"), ) -> ConversationMessagesResponse: """ Get all messages for a conversation belonging to an attack. @@ -359,7 +368,9 @@ async def get_conversation_messages( # pyrit-async-suffix-exempt 404: {"model": ProblemDetail, "description": "Attack not found"}, }, ) -async def get_conversations(attack_result_id: str) -> AttackConversationsResponse: # pyrit-async-suffix-exempt +async def get_conversations( + attack_result_id: IdentifierStr, +) -> AttackConversationsResponse: # pyrit-async-suffix-exempt """ Get all conversations belonging to an attack. @@ -391,7 +402,7 @@ async def get_conversations(attack_result_id: str) -> AttackConversationsRespons }, ) async def create_related_conversation( # pyrit-async-suffix-exempt - attack_result_id: str, + attack_result_id: IdentifierStr, request: CreateConversationRequest, ) -> CreateConversationResponse: """ @@ -434,7 +445,7 @@ async def create_related_conversation( # pyrit-async-suffix-exempt }, ) async def update_main_conversation( # pyrit-async-suffix-exempt - attack_result_id: str, + attack_result_id: IdentifierStr, request: UpdateMainConversationRequest, ) -> UpdateMainConversationResponse: """ @@ -479,7 +490,7 @@ async def update_main_conversation( # pyrit-async-suffix-exempt }, ) async def add_message( # pyrit-async-suffix-exempt - attack_result_id: str, + attack_result_id: IdentifierStr, request: AddMessageRequest, ) -> AddMessageResponse: """ diff --git a/pyrit/backend/routes/common.py b/pyrit/backend/routes/common.py index 146436cdc0..7462c5185a 100644 --- a/pyrit/backend/routes/common.py +++ b/pyrit/backend/routes/common.py @@ -3,6 +3,8 @@ """Shared route helpers.""" +from pyrit.backend.models.common import validate_label_filter + def parse_label_query_params(label_params: list[str] | None) -> dict[str, list[str]] | None: """ @@ -10,11 +12,12 @@ def parse_label_query_params(label_params: list[str] | None) -> dict[str, list[s Returns: dict[str, list[str]] | None: Labels grouped with OR-within-key semantics. + + Raises: + ValueError: If a label filter has no ``:`` separator or a part is too long. """ labels: dict[str, list[str]] = {} for param in label_params or []: - if ":" not in param: - continue - key, value = (part.strip() for part in param.split(":", 1)) - labels.setdefault(key, []).append(value) + key, _, value = validate_label_filter(param).partition(":") + labels.setdefault(key.strip(), []).append(value.strip()) return labels or None diff --git a/pyrit/backend/routes/configuration.py b/pyrit/backend/routes/configuration.py index a1cdb8f813..66f175457a 100644 --- a/pyrit/backend/routes/configuration.py +++ b/pyrit/backend/routes/configuration.py @@ -13,7 +13,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status from pyrit.backend.middleware.auth import AuthenticatedUser, require_admin -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import IdentifierStr, ProblemDetail from pyrit.backend.models.configuration import ( ConfigurationFileContent, EnvironmentFileContent, @@ -278,7 +278,7 @@ async def list_environment_files( # pyrit-async-suffix-exempt }, ) async def get_environment_file( # pyrit-async-suffix-exempt - file_id: str, + file_id: IdentifierStr, request: Request, ) -> EnvironmentFileContent: """ @@ -311,7 +311,7 @@ async def get_environment_file( # pyrit-async-suffix-exempt responses={404: {"model": ProblemDetail, "description": "Environment file not found"}}, ) async def update_environment_file( # pyrit-async-suffix-exempt - file_id: str, + file_id: IdentifierStr, body: UpdateEnvironmentFileRequest, request: Request, ) -> EnvironmentFileContent: diff --git a/pyrit/backend/routes/converters.py b/pyrit/backend/routes/converters.py index 83de572d5f..7423b1021f 100644 --- a/pyrit/backend/routes/converters.py +++ b/pyrit/backend/routes/converters.py @@ -10,7 +10,7 @@ from fastapi import APIRouter, HTTPException, status -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import IdentifierStr, ProblemDetail from pyrit.backend.models.converters import ( ConverterInstance, ConverterInstanceListResponse, @@ -97,7 +97,7 @@ async def create_converter(request: CreateConverterRequest) -> ConverterInstance 404: {"model": ProblemDetail, "description": "Converter not found"}, }, ) -async def get_converter(converter_id: str) -> ConverterInstance: # pyrit-async-suffix-exempt +async def get_converter(converter_id: IdentifierStr) -> ConverterInstance: # pyrit-async-suffix-exempt """ Get a converter instance by ID. @@ -123,7 +123,7 @@ async def get_converter(converter_id: str) -> ConverterInstance: # pyrit-async- 404: {"model": ProblemDetail, "description": "Converter not found"}, }, ) -async def delete_converter(converter_id: str) -> None: # pyrit-async-suffix-exempt +async def delete_converter(converter_id: IdentifierStr) -> None: # pyrit-async-suffix-exempt """Delete a converter instance by registry name.""" service = get_converter_service() if not await service.delete_converter_async(converter_id=converter_id): diff --git a/pyrit/backend/routes/initializers.py b/pyrit/backend/routes/initializers.py index 77ba39918f..6c99ca8067 100644 --- a/pyrit/backend/routes/initializers.py +++ b/pyrit/backend/routes/initializers.py @@ -18,7 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from pyrit.backend.middleware.auth import require_admin -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import CursorStr, IdentifierStr, ProblemDetail from pyrit.backend.models.initializers import ( ConfiguredInitializerSetting, CustomInitializerListResponse, @@ -88,7 +88,7 @@ def _check_custom_initializers_allowed(request: Request) -> None: ) async def list_initializers( # pyrit-async-suffix-exempt limit: int = Query(50, ge=1, le=200, description="Maximum items per page"), - cursor: str | None = Query(None, description="Pagination cursor (initializer_name to start after)"), + cursor: CursorStr | None = Query(None, description="Pagination cursor (initializer_name to start after)"), ) -> ListRegisteredInitializersResponse: """ List all available initializers. @@ -148,7 +148,7 @@ async def list_custom_initializers(request: Request) -> CustomInitializerListRes 404: {"model": ProblemDetail, "description": "Initializer not found"}, }, ) -async def get_initializer(initializer_name: str) -> RegisteredInitializer: # pyrit-async-suffix-exempt +async def get_initializer(initializer_name: IdentifierStr) -> RegisteredInitializer: # pyrit-async-suffix-exempt """ Get details for a specific initializer. @@ -223,7 +223,7 @@ async def register_initializer( # pyrit-async-suffix-exempt ) async def unregister_initializer( # pyrit-async-suffix-exempt request: Request, - initializer_name: str, + initializer_name: IdentifierStr, ) -> None: """ Remove a custom initializer from the registry. diff --git a/pyrit/backend/routes/labels.py b/pyrit/backend/routes/labels.py index d72ed505f9..8c514c3914 100644 --- a/pyrit/backend/routes/labels.py +++ b/pyrit/backend/routes/labels.py @@ -12,6 +12,7 @@ from fastapi import APIRouter, Query from pydantic import BaseModel, Field +from pyrit.backend.models.common import MAX_ITEMS, LabelFilterStr from pyrit.backend.routes.common import parse_label_query_params from pyrit.memory import CentralMemory @@ -39,13 +40,17 @@ async def get_label_options( # pyrit-async-suffix-exempt ), operator: list[Annotated[str, Field(max_length=128)]] | None = Query( None, + max_length=MAX_ITEMS, description="Narrow attack labels by operator.", ), operation: list[Annotated[str, Field(max_length=128)]] | None = Query( None, + max_length=MAX_ITEMS, description="Narrow attack labels by operation.", ), - label: list[str] | None = Query(None, description="Narrow attack labels by key:value filters."), + label: list[LabelFilterStr] | None = Query( + None, max_length=MAX_ITEMS, description="Narrow attack labels by key:value filters." + ), ) -> LabelOptionsResponse: """ Get unique label keys and values for filtering. diff --git a/pyrit/backend/routes/scenarios.py b/pyrit/backend/routes/scenarios.py index f73ccdb622..1b5ddbc766 100644 --- a/pyrit/backend/routes/scenarios.py +++ b/pyrit/backend/routes/scenarios.py @@ -14,7 +14,7 @@ from fastapi import APIRouter, HTTPException, Query, Request, status -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import MAX_ITEMS, CursorStr, IdentifierStr, LabelFilterStr, ProblemDetail from pyrit.backend.models.scenarios import ( ListRegisteredScenariosResponse, ScenarioRunListResponse, @@ -49,7 +49,7 @@ ) async def list_scenarios( # pyrit-async-suffix-exempt limit: int = Query(50, ge=1, le=200, description="Maximum items per page"), - cursor: str | None = Query(None, description="Pagination cursor (scenario_name to start after)"), + cursor: CursorStr | None = Query(None, description="Pagination cursor (scenario_name to start after)"), include_estimates: bool = Query(True, description="Wait for default run-size estimates"), ) -> ListRegisteredScenariosResponse: """ @@ -76,7 +76,7 @@ async def list_scenarios( # pyrit-async-suffix-exempt 404: {"model": ProblemDetail, "description": "Scenario not found"}, }, ) -async def get_scenario(scenario_name: str) -> RegisteredScenario: # pyrit-async-suffix-exempt +async def get_scenario(scenario_name: IdentifierStr) -> RegisteredScenario: # pyrit-async-suffix-exempt """ Get details for a specific scenario. @@ -108,7 +108,7 @@ async def get_scenario(scenario_name: str) -> RegisteredScenario: # pyrit-async ) async def estimate_scenario_run_size( # pyrit-async-suffix-exempt *, - scenario_name: str, + scenario_name: IdentifierStr, request: ScenarioRunSizeEstimateRequest, ) -> ScenarioRunSizeEstimate: """ @@ -193,7 +193,7 @@ async def start_scenario_run( # pyrit-async-suffix-exempt 409: {"model": ProblemDetail, "description": "Run is ineligible or has no saved launch configuration"}, }, ) -async def resume_scenario_run_async(*, scenario_result_id: str) -> ScenarioRunSummary: +async def resume_scenario_run_async(*, scenario_result_id: IdentifierStr) -> ScenarioRunSummary: """ Resume a failed run using its complete saved launch configuration. @@ -216,20 +216,23 @@ async def resume_scenario_run_async(*, scenario_result_id: str) -> ScenarioRunSu ) async def list_scenario_runs( # pyrit-async-suffix-exempt *, - scenario_names: list[str] | None = Query( + scenario_names: list[IdentifierStr] | None = Query( None, + max_length=MAX_ITEMS, description="Registered or persisted scenario names; repeated values are OR-matched.", ), run_statuses: list[ScenarioRunState] | None = Query( None, + max_length=MAX_ITEMS, description="Run states; repeated values are OR-matched.", ), - label: list[str] | None = Query( + label: list[LabelFilterStr] | None = Query( None, + max_length=MAX_ITEMS, description="key:value labels; OR within a key and AND across keys.", ), limit: int = Query(100, ge=1, le=100, description="Maximum items per page"), - cursor: str | None = Query(None, description="Opaque descending history cursor"), + cursor: CursorStr | None = Query(None, description="Opaque descending history cursor"), ) -> ScenarioRunListResponse: """ List tracked scenario runs (most recent first). @@ -275,7 +278,7 @@ async def get_scenario_run_queue() -> ScenarioQueueSnapshot: # pyrit-async-suff 404: {"model": ProblemDetail, "description": "Run not found"}, }, ) -async def get_scenario_run(scenario_result_id: str) -> ScenarioRunSummary: # pyrit-async-suffix-exempt +async def get_scenario_run(scenario_result_id: IdentifierStr) -> ScenarioRunSummary: # pyrit-async-suffix-exempt """ Get the current status and result of a scenario run. @@ -311,8 +314,8 @@ async def get_scenario_run(scenario_result_id: str) -> ScenarioRunSummary: # py ) async def get_scenario_run_progress( # pyrit-async-suffix-exempt *, - scenario_result_id: str, - since: str | None = Query(None, description="Opaque ascending progress cursor"), + scenario_result_id: IdentifierStr, + since: CursorStr | None = Query(None, description="Opaque ascending progress cursor"), limit: int = Query(100, ge=1, le=500), ) -> ScenarioRunProgress: """ @@ -350,7 +353,7 @@ async def get_scenario_run_progress( # pyrit-async-suffix-exempt 409: {"model": ProblemDetail, "description": "Run already in terminal state"}, }, ) -async def cancel_scenario_run(scenario_result_id: str) -> ScenarioRunSummary: # pyrit-async-suffix-exempt +async def cancel_scenario_run(scenario_result_id: IdentifierStr) -> ScenarioRunSummary: # pyrit-async-suffix-exempt """ Cancel a running scenario. @@ -382,7 +385,7 @@ async def cancel_scenario_run(scenario_result_id: str) -> ScenarioRunSummary: # 409: {"model": ProblemDetail, "description": "Run not yet completed"}, }, ) -async def get_scenario_run_results(scenario_result_id: str) -> ScenarioResult: # pyrit-async-suffix-exempt +async def get_scenario_run_results(scenario_result_id: IdentifierStr) -> ScenarioResult: # pyrit-async-suffix-exempt """ Get detailed results for a completed scenario run. diff --git a/pyrit/backend/routes/targets.py b/pyrit/backend/routes/targets.py index 03d41795a3..050b4ed0c7 100644 --- a/pyrit/backend/routes/targets.py +++ b/pyrit/backend/routes/targets.py @@ -10,7 +10,7 @@ from fastapi import APIRouter, HTTPException, Query, status -from pyrit.backend.models.common import ProblemDetail +from pyrit.backend.models.common import CursorStr, IdentifierStr, ProblemDetail from pyrit.backend.models.targets import ( CreateTargetRequest, TargetListResponse, @@ -31,7 +31,7 @@ ) async def list_targets( # pyrit-async-suffix-exempt limit: int = Query(50, ge=1, le=200, description="Maximum items per page"), - cursor: str | None = Query(None, description="Pagination cursor (target_registry_name)"), + cursor: CursorStr | None = Query(None, description="Pagination cursor (target_registry_name)"), ) -> TargetListResponse: """ List target instances with pagination. @@ -112,7 +112,7 @@ async def create_target( }, ) async def get_target( - target_registry_name: str, + target_registry_name: IdentifierStr, ) -> TargetInstance: # pyrit-async-suffix-exempt """ Get a target instance by registry name. diff --git a/pyrit/models/catalog/scenario.py b/pyrit/models/catalog/scenario.py index b8c234bd47..203d483780 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -16,7 +16,7 @@ from datetime import datetime from enum import Enum from math import prod -from typing import Any, Literal +from typing import Annotated, Any, Literal from pydantic import AliasChoices, BaseModel, Field, computed_field, field_validator, model_validator @@ -29,6 +29,29 @@ ScenarioDatasetSizeEstimate, ) +# Length and size limits for scenario run requests received over the REST API; the values +# match the backend request limits in ``pyrit.backend.models.common``. +_MAX_REQUEST_ITEMS = 100 +_MAX_REQUEST_NAME_LENGTH = 256 +_MAX_REQUEST_LABEL_KEY_LENGTH = 128 +_MAX_REQUEST_LABEL_VALUE_LENGTH = 1_024 +# Technique tokens can append converter modifiers (``technique:converter.:...``). +_MAX_REQUEST_TECHNIQUE_LENGTH = 4_096 +_RequestName = Annotated[str, Field(max_length=_MAX_REQUEST_NAME_LENGTH)] +_RequestNames = Annotated[list[_RequestName], Field(max_length=_MAX_REQUEST_ITEMS)] +_RequestTechniques = Annotated[ + list[Annotated[str, Field(max_length=_MAX_REQUEST_TECHNIQUE_LENGTH)]], Field(max_length=_MAX_REQUEST_ITEMS) +] +_RequestFilters = Annotated[dict[_RequestName, _RequestNames], Field(max_length=_MAX_REQUEST_ITEMS)] +_RequestParams = Annotated[dict[_RequestName, Any], Field(max_length=_MAX_REQUEST_ITEMS)] +_RequestLabels = Annotated[ + dict[ + Annotated[str, Field(max_length=_MAX_REQUEST_LABEL_KEY_LENGTH)], + Annotated[str, Field(max_length=_MAX_REQUEST_LABEL_VALUE_LENGTH)], + ], + Field(max_length=_MAX_REQUEST_ITEMS), +] + # Authoritative set of dataset seed filters exposed over the run request surface. Each entry # is used verbatim as a ``MemoryInterface.get_seeds`` keyword argument, so a filter key IS the # get_seeds kwarg. Every exposed filter must be a list-valued (Sequence) get_seeds parameter. @@ -359,23 +382,23 @@ class RegisteredScenario(BaseModel): class ScenarioRunSizeEstimateRequest(BaseModel): """Request-specific scenario run-size configuration.""" - adversarial_target_name: str | None = Field( + adversarial_target_name: _RequestName | None = Field( None, min_length=1, description="Registered multi-turn target overriding only the adversarial fallback for this request", ) - target_name: str | None = Field( + target_name: _RequestName | None = Field( None, description="Optional registered objective target used to resolve target-capability-dependent estimates", ) - techniques: list[str] | None = Field( + techniques: _RequestTechniques | None = Field( None, description="Technique names to estimate (uses scenario default if omitted)" ) - dataset_names: list[str] | None = Field( + dataset_names: _RequestNames | None = Field( None, description="Dataset names to estimate (uses scenario default if omitted)" ) max_dataset_size: int | None = Field(None, ge=1, description="Maximum selected logical seed groups") - dataset_filters: dict[str, list[str]] | None = Field( + dataset_filters: _RequestFilters | None = Field( None, description="Dataset seed filters keyed by field. Accepted keys: harm_categories, data_types.", ) @@ -383,7 +406,7 @@ class ScenarioRunSizeEstimateRequest(BaseModel): None, description="Override the scenario baseline default; forbidden scenarios reject true", ) - scenario_params: dict[str, Any] | None = Field( + scenario_params: _RequestParams | None = Field( None, description="Scenario-declared parameters such as Jailbreak template and attempt counts", ) @@ -403,20 +426,24 @@ def _validate_dataset_filters(cls, value: dict[str, list[str]] | None) -> dict[s class RunScenarioRequest(BaseModel): """Request body for starting a scenario run.""" - scenario_name: str = Field(..., description="Scenario name (e.g., 'foundry.red_team_agent')") - target_name: str = Field(..., description="Name of a registered target from the TargetRegistry") - adversarial_target_name: str | None = Field( + scenario_name: _RequestName = Field(..., description="Scenario name (e.g., 'foundry.red_team_agent')") + target_name: _RequestName = Field(..., description="Name of a registered target from the TargetRegistry") + adversarial_target_name: _RequestName | None = Field( None, min_length=1, description="Registered multi-turn target overriding only the adversarial fallback for this run", ) - initializers: list[str] | None = Field( + initializers: _RequestNames | None = Field( None, description="Initializer names to run before scenario (e.g., ['target', 'load_default_datasets'])" ) - techniques: list[str] | None = Field(None, description="Technique names to use (uses scenario default if omitted)") - dataset_names: list[str] | None = Field(None, description="Dataset names to use (uses scenario default if omitted)") + techniques: _RequestTechniques | None = Field( + None, description="Technique names to use (uses scenario default if omitted)" + ) + dataset_names: _RequestNames | None = Field( + None, description="Dataset names to use (uses scenario default if omitted)" + ) max_dataset_size: int | None = Field(None, ge=1, description="Maximum items per dataset") - dataset_filters: dict[str, list[str]] | None = Field( + dataset_filters: _RequestFilters | None = Field( None, description=( "Dataset seed filters keyed by field, applied before sampling. Accepted keys: harm_categories, data_types." @@ -427,19 +454,20 @@ class RunScenarioRequest(BaseModel): include_baseline: bool | None = Field( None, description="Override the scenario baseline default; forbidden scenarios reject true" ) - labels: dict[str, str] | None = Field(None, description="Labels to attach to memory entries") - scenario_params: dict[str, Any] | None = Field( + labels: _RequestLabels | None = Field(None, description="Labels to attach to memory entries") + scenario_params: _RequestParams | None = Field( None, description="Custom parameters for the scenario (passed to scenario.set_params_from_args). " "Keys are parameter names declared by the scenario's supported_parameters().", ) - initializer_args: dict[str, dict[str, Any]] | None = Field( + initializer_args: dict[_RequestName, _RequestParams] | None = Field( None, + max_length=_MAX_REQUEST_ITEMS, description="Per-initializer arguments keyed by initializer name. " "Each value is a dict of args passed to that initializer's set_params_from_args(). " "Example: {'target': {'endpoint': 'https://...'}}.", ) - scenario_result_id: str | None = Field( + scenario_result_id: _RequestName | None = Field( None, description="Optional ID of an existing ScenarioResult to resume. " "If provided, the scenario will resume from prior progress instead of starting fresh.", diff --git a/tests/unit/backend/test_api_routes.py b/tests/unit/backend/test_api_routes.py index e512b8e6cc..2fe2d9286d 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -28,7 +28,14 @@ MessageView, TargetResponseStatus, ) -from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.models.common import ( + MAX_CURSOR_LENGTH, + MAX_IDENTIFIER_LENGTH, + MAX_ITEMS, + MAX_LABEL_KEY_LENGTH, + MAX_LABEL_VALUE_LENGTH, + PaginationInfo, +) from pyrit.backend.models.converters import ( ConverterInstance, ConverterInstanceListResponse, @@ -41,6 +48,7 @@ TargetTypeResponse, ) from pyrit.backend.routes import version as version_routes +from pyrit.backend.routes.common import parse_label_query_params from pyrit.backend.routes.scores import _get_user_identifier from pyrit.backend.services.attack_service import AttackObjectiveConflictError from pyrit.backend.services.manual_send_scheduler import ManualSendConflictError, ManualSendQueueFullError @@ -789,27 +797,34 @@ def test_get_converter_options(self, client: TestClient) -> None: data = response.json() assert data["converter_types"] == ["Base64Converter", "ROT13Converter"] - def test_parse_labels_skips_param_without_colon(self, client: TestClient) -> None: - """Test that _parse_labels skips label params that have no colon.""" + def test_list_attacks_rejects_label_without_colon(self, client: TestClient) -> None: + """Test that a label filter without a key:value separator is rejected.""" with patch("pyrit.backend.routes.attacks.get_attack_service") as mock_get_service: mock_service = MagicMock() - mock_service.list_attacks_async = AsyncMock( - return_value=AttackListResponse( - items=[], - pagination=PaginationInfo(limit=20, has_more=False, next_cursor=None, prev_cursor=None), - ) - ) + mock_service.list_attacks_async = AsyncMock() mock_get_service.return_value = mock_service response = client.get("/api/attacks?label=nocolon&label=env:prod") - assert response.status_code == status.HTTP_200_OK - call_kwargs = mock_service.list_attacks_async.call_args[1] - # Only the valid label should be parsed - assert call_kwargs["labels"] == {"env": ["prod"]} + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT + assert "key:value" in response.json()["errors"][0]["message"] + mock_service.list_attacks_async.assert_not_called() + + @pytest.mark.parametrize("route", ["/api/attacks", "/api/labels", "/api/scenarios/runs"]) + @pytest.mark.parametrize( + "label", + ["k" * (MAX_LABEL_KEY_LENGTH + 1) + ":v", "k:" + "v" * (MAX_LABEL_VALUE_LENGTH + 1)], + ids=["key", "value"], + ) + def test_label_filter_part_over_limit_is_rejected(self, client: TestClient, route: str, label: str) -> None: + """Test that label filter keys and values each stay within the label limits.""" + response = client.get(route, params={"label": label}) + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT - def test_parse_labels_all_invalid_returns_none(self, client: TestClient) -> None: - """Test that _parse_labels returns None when all params lack colons.""" + def test_label_filter_at_limits_is_accepted(self, client: TestClient) -> None: + """Test that padded label filters at the key and value limits are parsed.""" + key, value = "k" * MAX_LABEL_KEY_LENGTH, "v" * MAX_LABEL_VALUE_LENGTH with patch("pyrit.backend.routes.attacks.get_attack_service") as mock_get_service: mock_service = MagicMock() mock_service.list_attacks_async = AsyncMock( @@ -820,11 +835,40 @@ def test_parse_labels_all_invalid_returns_none(self, client: TestClient) -> None ) mock_get_service.return_value = mock_service - response = client.get("/api/attacks?label=nocolon&label=alsonocolon") + response = client.get("/api/attacks", params={"label": f" {key} : {value} "}) assert response.status_code == status.HTTP_200_OK - call_kwargs = mock_service.list_attacks_async.call_args[1] - assert call_kwargs["labels"] is None + assert mock_service.list_attacks_async.call_args[1]["labels"] == {key: [value]} + + def test_parse_label_query_params_rejects_missing_separator(self) -> None: + """Test that the shared label parser never drops a malformed filter.""" + with pytest.raises(ValueError, match="key:value"): + parse_label_query_params(["env:prod", "nocolon"]) + + @pytest.mark.parametrize( + "query", + [ + "label=" + "&label=".join(f"k{i}:v" for i in range(MAX_ITEMS + 1)), + "cursor=" + "c" * (MAX_CURSOR_LENGTH + 1), + "attack_types=" + "a" * (MAX_IDENTIFIER_LENGTH + 1), + "min_turns=100000000000000000000", + ], + ) + def test_list_attacks_rejects_oversized_query_values(self, client: TestClient, query: str) -> None: + """Test that oversized filter values are rejected before the service runs.""" + with patch("pyrit.backend.routes.attacks.get_attack_service") as mock_get_service: + response = client.get(f"/api/attacks?{query}") + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT + mock_get_service.assert_not_called() + + def test_get_attack_rejects_oversized_id(self, client: TestClient) -> None: + """Test that an oversized path identifier is rejected before the service runs.""" + with patch("pyrit.backend.routes.attacks.get_attack_service") as mock_get_service: + response = client.get(f"/api/attacks/{'a' * (MAX_IDENTIFIER_LENGTH + 1)}") + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT + mock_get_service.assert_not_called() def test_parse_labels_value_with_extra_colons(self, client: TestClient) -> None: """Test that _parse_labels handles values containing colons (split on first only).""" diff --git a/tests/unit/backend/test_common_models.py b/tests/unit/backend/test_common_models.py index 00ae1b794f..13e37c1d72 100644 --- a/tests/unit/backend/test_common_models.py +++ b/tests/unit/backend/test_common_models.py @@ -5,12 +5,40 @@ Tests for backend common models. """ +import uuid + +import pytest +from pydantic import BaseModel, TypeAdapter, ValidationError + +from pyrit.backend.models.attacks import ( + AddMessageRequest, + CreateAttackRequest, + MessagePieceRequest, + UpdateAttackRequest, +) from pyrit.backend.models.common import ( + MAX_FILE_CONTENT_LENGTH, + MAX_IDENTIFIER_LENGTH, + MAX_ITEMS, + MAX_LABEL_KEY_LENGTH, + MAX_LABEL_VALUE_LENGTH, + MAX_TEXT_LENGTH, FieldError, + LabelFilterStr, PaginationInfo, ProblemDetail, filter_sensitive_fields, ) +from pyrit.backend.models.configuration import ( + ReinitializeRequest, + UpdateConfigurationFileRequest, + UpdateEnvironmentFileRequest, +) +from pyrit.backend.models.converters import ConverterPreviewRequest, CreateConverterRequest +from pyrit.backend.models.initializers import RegisterInitializerRequest +from pyrit.backend.models.scores import ManualScoreRequest +from pyrit.backend.models.targets import CreateTargetRequest +from pyrit.models.catalog import scenario as scenario_catalog class TestPaginationInfo: @@ -371,3 +399,112 @@ def test_problem_detail_serialization(self) -> None: assert "instance" not in data # None should be excluded assert data["type"] == "/errors/test" + + +@pytest.mark.parametrize( + "overrides", + [ + {"target_registry_name": "t" * (MAX_IDENTIFIER_LENGTH + 1)}, + {"source_conversation_id": "c" * (MAX_IDENTIFIER_LENGTH + 1)}, + {"name": "n" * (MAX_TEXT_LENGTH + 1)}, + {"labels": {f"key{i}": "value" for i in range(MAX_ITEMS + 1)}}, + {"labels": {"k" * (MAX_LABEL_KEY_LENGTH + 1): "value"}}, + {"labels": {"key": "v" * (MAX_LABEL_VALUE_LENGTH + 1)}}, + {"cutoff_index": -1}, + ], +) +def test_create_attack_request_rejects_values_over_limits(overrides: dict[str, object]) -> None: + with pytest.raises(ValidationError): + CreateAttackRequest.model_validate({"target_registry_name": "target", **overrides}) + + +def test_create_attack_request_accepts_values_at_limits() -> None: + request = CreateAttackRequest( + target_registry_name="t" * MAX_IDENTIFIER_LENGTH, + name="n" * MAX_TEXT_LENGTH, + labels={f"{i:0{MAX_LABEL_KEY_LENGTH}d}": "v" * MAX_LABEL_VALUE_LENGTH for i in range(MAX_ITEMS)}, + cutoff_index=0, + ) + + assert len(request.labels or {}) == MAX_ITEMS + + +def test_update_attack_request_rejects_oversized_objective() -> None: + with pytest.raises(ValidationError): + UpdateAttackRequest(objective="o" * (MAX_TEXT_LENGTH + 1)) + + +@pytest.mark.parametrize( + "overrides", + [ + {"target_conversation_id": "c" * (MAX_IDENTIFIER_LENGTH + 1)}, + {"converter_ids": ["converter"] * (MAX_ITEMS + 1)}, + {"request_converter_configurations": [{"converter_ids": ["converter"]}] * (MAX_ITEMS + 1)}, + {"pieces": [{"original_value": "text", "mime_type": "m" * (MAX_IDENTIFIER_LENGTH + 1)}]}, + ], +) +def test_add_message_request_rejects_values_over_limits(overrides: dict[str, object]) -> None: + payload = {"pieces": [{"original_value": "text"}], "target_conversation_id": "conversation", **overrides} + + with pytest.raises(ValidationError): + AddMessageRequest.model_validate(payload) + + +def test_message_piece_content_is_not_length_limited() -> None: + piece = MessagePieceRequest(original_value="x" * (MAX_TEXT_LENGTH + 1)) + + assert len(piece.original_value) == MAX_TEXT_LENGTH + 1 + + +@pytest.mark.parametrize( + ("model", "payload", "field"), + [ + (UpdateConfigurationFileRequest, {"content": "x" * (MAX_FILE_CONTENT_LENGTH + 1), "version": "v"}, "content"), + (UpdateEnvironmentFileRequest, {"content": "x" * (MAX_FILE_CONTENT_LENGTH + 1), "version": "v"}, "content"), + (ReinitializeRequest, {"version": "v" * (MAX_IDENTIFIER_LENGTH + 1)}, "version"), + ( + RegisterInitializerRequest, + {"name": "custom", "script_content": "x" * (MAX_FILE_CONTENT_LENGTH + 1)}, + "script_content", + ), + ( + CreateConverterRequest, + {"name": "c", "type": "T", "params": {f"p{i}": 1 for i in range(MAX_ITEMS + 1)}}, + "params", + ), + (ConverterPreviewRequest, {"original_value": "x", "converter_ids": ["c"] * (MAX_ITEMS + 1)}, "converter_ids"), + (CreateTargetRequest, {"type": "t" * (MAX_IDENTIFIER_LENGTH + 1)}, "type"), + ( + ManualScoreRequest, + { + "attack_result_id": str(uuid.uuid4()), + "message_id": str(uuid.uuid4()), + "value": True, + "rationale": "r" * (MAX_TEXT_LENGTH + 1), + }, + "rationale", + ), + ], +) +def test_request_models_reject_values_over_limits( + model: type[BaseModel], payload: dict[str, object], field: str +) -> None: + with pytest.raises(ValidationError) as error: + model.model_validate(payload) + + assert [item["loc"][0] for item in error.value.errors()] == [field] + + +def test_scenario_request_limits_match_backend_limits() -> None: + assert scenario_catalog._MAX_REQUEST_ITEMS == MAX_ITEMS + assert scenario_catalog._MAX_REQUEST_NAME_LENGTH == MAX_IDENTIFIER_LENGTH + assert scenario_catalog._MAX_REQUEST_LABEL_KEY_LENGTH == MAX_LABEL_KEY_LENGTH + assert scenario_catalog._MAX_REQUEST_LABEL_VALUE_LENGTH == MAX_LABEL_VALUE_LENGTH + + +def test_label_filter_limits_ignore_surrounding_whitespace() -> None: + padded = f" {'k' * MAX_LABEL_KEY_LENGTH} : {'v' * MAX_LABEL_VALUE_LENGTH} " + + assert TypeAdapter(LabelFilterStr).validate_python(padded) == padded + with pytest.raises(ValidationError, match="limited"): + TypeAdapter(LabelFilterStr).validate_python(f"{'k' * (MAX_LABEL_KEY_LENGTH + 1)}:v") diff --git a/tests/unit/backend/test_request_size.py b/tests/unit/backend/test_request_size.py new file mode 100644 index 0000000000..ed38e43110 --- /dev/null +++ b/tests/unit/backend/test_request_size.py @@ -0,0 +1,93 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from collections.abc import Iterator +from unittest.mock import patch + +import pytest +from fastapi import FastAPI, Request +from fastapi.testclient import TestClient + +from pyrit.backend.main import app as backend_app +from pyrit.backend.middleware.error_handlers import register_error_handlers +from pyrit.backend.middleware.request_size import RequestSizeLimitMiddleware + +_LIMIT = 16 +_TOO_LARGE = { + "type": "/errors/request-too-large", + "title": "Content Too Large", + "status": 413, + "detail": "The request is larger than the backend accepts.", +} + + +@pytest.fixture +def client() -> Iterator[TestClient]: + app = FastAPI() + register_error_handlers(app) + app.add_middleware(RequestSizeLimitMiddleware) + + @app.post("/echo") + async def echo(request: Request) -> dict[str, int]: + return {"size": len(await request.body())} + + @app.get("/items") + async def items(q: str = "") -> dict[str, int]: + return {"size": len(q)} + + with ( + patch.object(RequestSizeLimitMiddleware, "MAX_BODY_BYTES", _LIMIT), + patch.object(RequestSizeLimitMiddleware, "MAX_URL_LENGTH", 64), + TestClient(app) as test_client, + ): + yield test_client + + +def test_body_within_limit_reaches_handler(client: TestClient) -> None: + response = client.post("/echo", content=b"a" * _LIMIT) + + assert response.status_code == 200 + assert response.json() == {"size": _LIMIT} + + +@pytest.mark.parametrize( + "content", + [b"a" * (_LIMIT + 1), iter([b"a" * _LIMIT, b"b"])], + ids=["declared", "streamed"], +) +def test_body_over_limit_returns_problem(client: TestClient, content: bytes | Iterator[bytes]) -> None: + response = client.post("/echo", content=content) + + assert response.status_code == 413 + assert response.headers["content-type"] == "application/problem+json" + assert response.json() == _TOO_LARGE + + +def test_long_url_returns_414(client: TestClient) -> None: + response = client.get("/items", params={"q": "x" * 64}) + + assert response.status_code == 414 + assert response.headers["content-type"] == "application/problem+json" + assert response.json() == {**_TOO_LARGE, "title": "URI Too Long", "status": 414} + + +def test_url_within_limit_reaches_handler(client: TestClient) -> None: + response = client.get("/items", params={"q": "x" * 8}) + + assert response.status_code == 200 + assert response.json() == {"size": 8} + + +def test_backend_app_registers_middleware() -> None: + assert RequestSizeLimitMiddleware in [middleware.cls for middleware in backend_app.user_middleware] + + +def test_backend_app_returns_413_for_streamed_body_over_limit(compatibility_headers: dict[str, str]) -> None: + client = TestClient(backend_app, headers={**compatibility_headers, "Content-Type": "application/json"}) + + with patch.object(RequestSizeLimitMiddleware, "MAX_BODY_BYTES", _LIMIT): + response = client.post("/api/attacks", content=iter([b"{" * _LIMIT, b"}"])) + + assert response.status_code == 413 + assert response.headers["content-type"] == "application/problem+json" + assert response.json() == _TOO_LARGE diff --git a/tests/unit/models/test_scenario_request.py b/tests/unit/models/test_scenario_request.py index 938ed2ef04..4943382eea 100644 --- a/tests/unit/models/test_scenario_request.py +++ b/tests/unit/models/test_scenario_request.py @@ -6,7 +6,7 @@ import pytest from pydantic import ValidationError -from pyrit.models.catalog.scenario import DATASET_FILTERS, RunScenarioRequest +from pyrit.models.catalog.scenario import DATASET_FILTERS, RunScenarioRequest, ScenarioRunSizeEstimateRequest def _make_request(*, dataset_filters: dict[str, list[str]] | None) -> RunScenarioRequest: @@ -64,3 +64,62 @@ def _allows_sequence(annotation: object) -> bool: for name in DATASET_FILTERS: assert name in hints, f"'{name}' is not a MemoryInterface.get_seeds_async parameter" assert _allows_sequence(hints[name]), f"'{name}' must be a Sequence-typed get_seeds parameter" + + +@pytest.mark.parametrize( + "overrides", + [ + {"scenario_name": "s" * 257}, + {"scenario_result_id": "r" * 257}, + {"techniques": ["technique"] * 101}, + {"initializers": ["i" * 257]}, + {"dataset_filters": {"harm_categories": ["cyber"] * 101}}, + {"labels": {f"key{i}": "value" for i in range(101)}}, + {"labels": {"k" * 129: "value"}}, + {"labels": {"key": "v" * 1_025}}, + {"scenario_params": {f"param{i}": 1 for i in range(101)}}, + {"initializer_args": {"target": {f"arg{i}": 1 for i in range(101)}}}, + ], +) +def test_run_request_rejects_values_over_limits(overrides: dict[str, object]) -> None: + with pytest.raises(ValidationError): + RunScenarioRequest.model_validate({"scenario_name": "s", "target_name": "t", **overrides}) + + +def test_run_request_accepts_values_at_limits() -> None: + request = RunScenarioRequest( + scenario_name="s" * 256, + target_name="t", + techniques=["technique"] * 100, + labels={"k" * 128: "v" * 1_024}, + ) + + assert request.techniques is not None + assert len(request.techniques) == 100 + + +@pytest.mark.parametrize("model", [RunScenarioRequest, ScenarioRunSizeEstimateRequest]) +def test_requests_accept_technique_with_converter_modifiers(model: type) -> None: + technique = "prompt_sending" + "".join(f":converter.{'c' * 64}" for _ in range(4)) + + request = model.model_validate({"scenario_name": "s", "target_name": "t", "techniques": [technique]}) + + assert request.techniques == [technique] + + +@pytest.mark.parametrize("model", [RunScenarioRequest, ScenarioRunSizeEstimateRequest]) +def test_requests_reject_oversized_technique(model: type) -> None: + with pytest.raises(ValidationError): + model.model_validate({"scenario_name": "s", "target_name": "t", "techniques": ["t" * 4_097]}) + + +def test_estimate_request_rejects_too_many_dataset_names() -> None: + with pytest.raises(ValidationError): + ScenarioRunSizeEstimateRequest(dataset_names=["dataset"] * 101) + + +def test_run_request_bounds_dataset_filter_keys() -> None: + with pytest.raises(ValidationError) as error: + RunScenarioRequest(scenario_name="s", target_name="t", dataset_filters={"k" * 257: ["x"]}) + + assert error.value.errors()[0]["type"] == "string_too_long" From 7c7afc81fbcb85649e3e2d932d9177b5f072c678 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Fri, 2 Oct 2026 19:36:08 +0000 Subject: [PATCH 2/3] MAINT: Share request field limits from pyrit.models The scenario request models duplicated the backend's identifier, list, and label limits because pyrit.models cannot import from the backend, and a test kept the two copies equal. The shared limits now live in pyrit.models.request_limits; the scenario models and the backend import them. HTTP body and URL limits stay in the backend. --- pyrit/backend/models/common.py | 13 ++++++------ pyrit/models/catalog/scenario.py | 25 ++++++++++-------------- pyrit/models/request_limits.py | 9 +++++++++ tests/unit/backend/test_common_models.py | 8 -------- 4 files changed, 25 insertions(+), 30 deletions(-) create mode 100644 pyrit/models/request_limits.py diff --git a/pyrit/backend/models/common.py b/pyrit/backend/models/common.py index 2a18939f92..d58dbd5ba3 100644 --- a/pyrit/backend/models/common.py +++ b/pyrit/backend/models/common.py @@ -11,17 +11,16 @@ from pydantic import AfterValidator, BaseModel, Field +from pyrit.models.request_limits import MAX_IDENTIFIER_LENGTH, MAX_ITEMS, MAX_LABEL_KEY_LENGTH, MAX_LABEL_VALUE_LENGTH + REGISTRY_INSTANCE_NAME_PATTERN = r"^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$" -# Request limits. Prompt content (message pieces, system prompts, preview input) and free-form -# values (metadata and parameter values) are only limited by the request body size, so long -# prompts and base64 media keep working. -MAX_IDENTIFIER_LENGTH = 256 +# Request limits. Identifier, list, and label limits are shared with ``pyrit.models``. Prompt +# content (message pieces, system prompts, preview input) and free-form values (metadata and +# parameter values) are only limited by the request body size, so long prompts and base64 media +# keep working. MAX_CURSOR_LENGTH = 1_024 MAX_TEXT_LENGTH = 100_000 -MAX_ITEMS = 100 -MAX_LABEL_KEY_LENGTH = 128 -MAX_LABEL_VALUE_LENGTH = 1_024 MAX_FILE_CONTENT_LENGTH = 1_048_576 IdentifierStr = Annotated[str, Field(max_length=MAX_IDENTIFIER_LENGTH)] diff --git a/pyrit/models/catalog/scenario.py b/pyrit/models/catalog/scenario.py index 203d483780..7e9bde5aaa 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -21,6 +21,7 @@ from pydantic import AliasChoices, BaseModel, Field, computed_field, field_validator, model_validator from pyrit.models.parameter import Parameter +from pyrit.models.request_limits import MAX_IDENTIFIER_LENGTH, MAX_ITEMS, MAX_LABEL_KEY_LENGTH, MAX_LABEL_VALUE_LENGTH from pyrit.models.results.scenario_result import ScenarioRunState from pyrit.models.retry_event import RetryEvent from pyrit.models.scenario_dataset_size_estimate import ( @@ -29,27 +30,21 @@ ScenarioDatasetSizeEstimate, ) -# Length and size limits for scenario run requests received over the REST API; the values -# match the backend request limits in ``pyrit.backend.models.common``. -_MAX_REQUEST_ITEMS = 100 -_MAX_REQUEST_NAME_LENGTH = 256 -_MAX_REQUEST_LABEL_KEY_LENGTH = 128 -_MAX_REQUEST_LABEL_VALUE_LENGTH = 1_024 # Technique tokens can append converter modifiers (``technique:converter.:...``). _MAX_REQUEST_TECHNIQUE_LENGTH = 4_096 -_RequestName = Annotated[str, Field(max_length=_MAX_REQUEST_NAME_LENGTH)] -_RequestNames = Annotated[list[_RequestName], Field(max_length=_MAX_REQUEST_ITEMS)] +_RequestName = Annotated[str, Field(max_length=MAX_IDENTIFIER_LENGTH)] +_RequestNames = Annotated[list[_RequestName], Field(max_length=MAX_ITEMS)] _RequestTechniques = Annotated[ - list[Annotated[str, Field(max_length=_MAX_REQUEST_TECHNIQUE_LENGTH)]], Field(max_length=_MAX_REQUEST_ITEMS) + list[Annotated[str, Field(max_length=_MAX_REQUEST_TECHNIQUE_LENGTH)]], Field(max_length=MAX_ITEMS) ] -_RequestFilters = Annotated[dict[_RequestName, _RequestNames], Field(max_length=_MAX_REQUEST_ITEMS)] -_RequestParams = Annotated[dict[_RequestName, Any], Field(max_length=_MAX_REQUEST_ITEMS)] +_RequestFilters = Annotated[dict[_RequestName, _RequestNames], Field(max_length=MAX_ITEMS)] +_RequestParams = Annotated[dict[_RequestName, Any], Field(max_length=MAX_ITEMS)] _RequestLabels = Annotated[ dict[ - Annotated[str, Field(max_length=_MAX_REQUEST_LABEL_KEY_LENGTH)], - Annotated[str, Field(max_length=_MAX_REQUEST_LABEL_VALUE_LENGTH)], + Annotated[str, Field(max_length=MAX_LABEL_KEY_LENGTH)], + Annotated[str, Field(max_length=MAX_LABEL_VALUE_LENGTH)], ], - Field(max_length=_MAX_REQUEST_ITEMS), + Field(max_length=MAX_ITEMS), ] # Authoritative set of dataset seed filters exposed over the run request surface. Each entry @@ -462,7 +457,7 @@ class RunScenarioRequest(BaseModel): ) initializer_args: dict[_RequestName, _RequestParams] | None = Field( None, - max_length=_MAX_REQUEST_ITEMS, + max_length=MAX_ITEMS, description="Per-initializer arguments keyed by initializer name. " "Each value is a dict of args passed to that initializer's set_params_from_args(). " "Example: {'target': {'endpoint': 'https://...'}}.", diff --git a/pyrit/models/request_limits.py b/pyrit/models/request_limits.py new file mode 100644 index 0000000000..c9a20196ce --- /dev/null +++ b/pyrit/models/request_limits.py @@ -0,0 +1,9 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Length and size limits for request values that more than one API layer validates.""" + +MAX_ITEMS = 100 +MAX_IDENTIFIER_LENGTH = 256 +MAX_LABEL_KEY_LENGTH = 128 +MAX_LABEL_VALUE_LENGTH = 1_024 diff --git a/tests/unit/backend/test_common_models.py b/tests/unit/backend/test_common_models.py index 13e37c1d72..5f1a65878a 100644 --- a/tests/unit/backend/test_common_models.py +++ b/tests/unit/backend/test_common_models.py @@ -38,7 +38,6 @@ from pyrit.backend.models.initializers import RegisterInitializerRequest from pyrit.backend.models.scores import ManualScoreRequest from pyrit.backend.models.targets import CreateTargetRequest -from pyrit.models.catalog import scenario as scenario_catalog class TestPaginationInfo: @@ -495,13 +494,6 @@ def test_request_models_reject_values_over_limits( assert [item["loc"][0] for item in error.value.errors()] == [field] -def test_scenario_request_limits_match_backend_limits() -> None: - assert scenario_catalog._MAX_REQUEST_ITEMS == MAX_ITEMS - assert scenario_catalog._MAX_REQUEST_NAME_LENGTH == MAX_IDENTIFIER_LENGTH - assert scenario_catalog._MAX_REQUEST_LABEL_KEY_LENGTH == MAX_LABEL_KEY_LENGTH - assert scenario_catalog._MAX_REQUEST_LABEL_VALUE_LENGTH == MAX_LABEL_VALUE_LENGTH - - def test_label_filter_limits_ignore_surrounding_whitespace() -> None: padded = f" {'k' * MAX_LABEL_KEY_LENGTH} : {'v' * MAX_LABEL_VALUE_LENGTH} " From d276a44884d292fc43347a4b6a5a0b31be496230 Mon Sep 17 00:00:00 2001 From: varunj-msft Date: Mon, 5 Oct 2026 23:54:15 +0000 Subject: [PATCH 3/3] FIX: Leave expected_objective unbounded so long objectives stay editable --- pyrit/backend/models/attacks.py | 2 +- tests/unit/backend/test_common_models.py | 6 ++++++ 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index e54d57c2c8..dc7f34d210 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -521,7 +521,7 @@ class UpdateAttackRequest(BaseModel): description="Updated attack outcome", ) objective: TextStr | None = Field(default=None, description="Shared objective for all conversations in the attack") - expected_objective: TextStr | None = Field( + expected_objective: str | None = Field( default=None, description="Objective read before editing, for conflict detection" ) diff --git a/tests/unit/backend/test_common_models.py b/tests/unit/backend/test_common_models.py index 5f1a65878a..10c062fec8 100644 --- a/tests/unit/backend/test_common_models.py +++ b/tests/unit/backend/test_common_models.py @@ -433,6 +433,12 @@ def test_update_attack_request_rejects_oversized_objective() -> None: UpdateAttackRequest(objective="o" * (MAX_TEXT_LENGTH + 1)) +def test_update_attack_request_expected_objective_is_not_length_limited() -> None: + request = UpdateAttackRequest(objective="o", expected_objective="o" * (MAX_TEXT_LENGTH + 1)) + + assert len(request.expected_objective or "") == MAX_TEXT_LENGTH + 1 + + @pytest.mark.parametrize( "overrides", [