Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions pyrit/backend/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,16 @@ from `frontend`. Set `PYRIT_PYTHON` to this worktree's Python interpreter if Vit
find `python` (for example, `<worktree>\.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:
Expand Down
4 changes: 4 additions & 0 deletions pyrit/backend/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
14 changes: 14 additions & 0 deletions pyrit/backend/middleware/error_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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,
Expand Down
88 changes: 88 additions & 0 deletions pyrit/backend/middleware/request_size.py
Original file line number Diff line number Diff line change
@@ -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)
54 changes: 36 additions & 18 deletions pyrit/backend/models/attacks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.",
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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. "
Expand Down Expand Up @@ -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"
)
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand All @@ -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.",
)

Expand All @@ -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.",
Expand Down
53 changes: 49 additions & 4 deletions pyrit/backend/models/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
Loading
Loading