diff --git a/pyrit/backend/README.md b/pyrit/backend/README.md index ab95f96d7b..9833855ea4 100644 --- a/pyrit/backend/README.md +++ b/pyrit/backend/README.md @@ -125,6 +125,16 @@ from `frontend`. Set `PYRIT_PYTHON` to this worktree's Python interpreter if Vit find `python` (for example, `\.venv\Scripts\python.exe` on Windows). Stop these test-owned servers after the run. +### 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 77bbf75916..25fd7e0998 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, @@ -92,6 +93,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 da6b286e85..dc7f34d210 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, @@ -372,17 +378,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.", @@ -426,7 +434,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 @@ -472,14 +480,16 @@ 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 | None = Field( + name: TextStr | None = Field(None, description="Attack name/label") + target_registry_name: IdentifierStr | None = Field( None, description="Target registry name, or None for a saved unbound attack" ) - source_conversation_id: str | None = Field( + 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. " @@ -510,7 +520,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") expected_objective: str | None = Field( default=None, description="Objective read before editing, for conflict detection" ) @@ -563,8 +573,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): @@ -614,7 +626,7 @@ def _validate_destination(self) -> "SaveConversationRequest": 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): @@ -633,19 +645,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.", ) @@ -664,25 +679,28 @@ class AddMessageRequest(MessageRequest): 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..d58dbd5ba3 100644 --- a/pyrit/backend/models/common.py +++ b/pyrit/backend/models/common.py @@ -2,17 +2,62 @@ # 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 + +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. 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_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 c75157de9c..bb75f1f33b 100644 --- a/pyrit/backend/routes/attacks.py +++ b/pyrit/backend/routes/attacks.py @@ -32,7 +32,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 @@ -70,14 +76,16 @@ async def save_conversation_async(*, body: SaveConversationRequest, request: Req 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). " @@ -102,21 +110,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. " @@ -255,7 +264,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. @@ -285,7 +294,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: """ @@ -316,7 +325,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. @@ -342,8 +351,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. @@ -382,7 +391,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. @@ -414,7 +425,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: """ @@ -457,7 +468,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: """ @@ -502,7 +513,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 19029d665c..70966e8a40 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. @@ -305,8 +308,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: """ @@ -340,7 +343,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. @@ -372,7 +375,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..7e9bde5aaa 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -16,11 +16,12 @@ 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 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,6 +30,23 @@ ScenarioDatasetSizeEstimate, ) +# Technique tokens can append converter modifiers (``technique:converter.:...``). +_MAX_REQUEST_TECHNIQUE_LENGTH = 4_096 +_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_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_LABEL_KEY_LENGTH)], + Annotated[str, Field(max_length=MAX_LABEL_VALUE_LENGTH)], + ], + Field(max_length=MAX_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 +377,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 +401,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 +421,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 +449,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_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/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_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..10c062fec8 100644 --- a/tests/unit/backend/test_common_models.py +++ b/tests/unit/backend/test_common_models.py @@ -5,12 +5,39 @@ 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 class TestPaginationInfo: @@ -371,3 +398,111 @@ 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)) + + +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", + [ + {"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_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"