Skip to content
Open
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
2 changes: 1 addition & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ pre-commit:
pre-commit run --all-files

ty:
$(CMD) ty check $(PYMODULE) $(UNIT_TESTS)
uv run --frozen --extra all --link-mode=copy -m ty check $(PYMODULE) $(UNIT_TESTS)

# Build the full documentation site:
# 1. Generate API reference JSON from Python source (griffe)
Expand Down
25 changes: 23 additions & 2 deletions doc/getting_started/install_local_dev.md
Original file line number Diff line number Diff line change
Expand Up @@ -202,10 +202,31 @@ uv run ruff check --fix .

#### Running Type Checker

```bash
uv run ty check pyrit/
Run checks from the repository root using its own uv environment. Install all optional
dependencies so guarded imports have the same dependency coverage as the production
pre-commit hook:

```powershell
uv sync --frozen --extra all
uv run --frozen --no-sync ty check pyrit
uv run --frozen --no-sync ty check pyrit tests\unit
```

The first check covers production code, matching the CI typing hook. The second also
checks unit tests and is the scope of `make ty`. Pytest and ty both resolve test helpers
from the `tests` directory; use tier-root imports such as `from unit.mocks import MockPromptTarget`.

Checking all test tiers and build scripts is a separate, wider diagnostic scope:

```powershell
uv run --frozen --no-sync ty check pyrit tests build_scripts
```

That wider scope is not the CI typing hook and may still report errors outside unit tests.
The lock file pins the checker version. Record the environment's Python version separately
from ty's target version, which defaults to the minimum supported Python version unless
overridden with `--python-version`.

#### Pre-commit Hooks

```bash
Expand Down
5 changes: 4 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -178,7 +178,7 @@ pyrit_shell = "pyrit.cli.pyrit_shell:main"
addopts = [
"--import-mode=importlib",
]
pythonpath = ["."]
pythonpath = [".", "tests"]
asyncio_default_fixture_loop_scope = "function"
asyncio_mode = "auto"
filterwarnings = [
Expand All @@ -193,6 +193,9 @@ filterwarnings = [
]

[tool.ty]
[tool.ty.environment]
extra-paths = ["tests"]

[tool.ty.rules]
all = "error"
# Most rules under `all = "error"` are already clean for pyrit/. A few remain
Expand Down
7 changes: 4 additions & 3 deletions pyrit/exceptions/exception_classes.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import uuid
from abc import ABC
from collections.abc import Callable, Sequence
from typing import Any
from typing import Any, TypeVar

from openai import RateLimitError
from tenacity import (
Expand All @@ -25,6 +25,7 @@
from pyrit.models import Message, MessagePiece, construct_response_from_request

logger = logging.getLogger(__name__)
_WrappedCallable = TypeVar("_WrappedCallable", bound=Callable[..., Any])


def _get_custom_result_retry_max_num_attempts() -> int:
Expand Down Expand Up @@ -345,7 +346,7 @@ class ExperimentalWarning(FutureWarning):

def pyrit_custom_result_retry(
retry_function: Callable[..., bool], retry_max_num_attempts: int | None = None
) -> Callable[..., Any]:
) -> Callable[[_WrappedCallable], _WrappedCallable]:
"""
Apply retry logic with exponential backoff to a function.

Expand All @@ -364,7 +365,7 @@ def pyrit_custom_result_retry(

"""

def inner_retry(func: Callable[..., Any]) -> Callable[..., Any]:
def inner_retry(func: _WrappedCallable) -> _WrappedCallable:
# Use static value if explicitly provided, otherwise use dynamic getter
stop_strategy: stop_base
if retry_max_num_attempts is not None:
Expand Down
4 changes: 2 additions & 2 deletions pyrit/prompt_target/a2a_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from collections.abc import AsyncGenerator
from dataclasses import dataclass
from email.utils import parsedate_to_datetime
from typing import TYPE_CHECKING, Any, Literal, cast
from typing import TYPE_CHECKING, Any, Literal
from weakref import WeakValueDictionary

import httpx
Expand Down Expand Up @@ -341,7 +341,7 @@ async def _await_task_async(self, *, client: Client, task: a2a_pb2.Task) -> a2a_
while task.status.state in (a2a_pb2.TASK_STATE_SUBMITTED, a2a_pb2.TASK_STATE_WORKING):
await asyncio.sleep(delay)
try:
task = cast("a2a_pb2.Task", await self._get_task_async(client=client, task_id=task.id))
task = await self._get_task_async(client=client, task_id=task.id)
delay = self._poll_interval_seconds
except (A2AError, httpx.HTTPError) as exc:
response = self._rate_limit_response(exc)
Expand Down
13 changes: 9 additions & 4 deletions pyrit/prompt_target/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,9 @@

import asyncio
import logging
from collections.abc import Callable
from collections.abc import Awaitable, Callable, Coroutine
from functools import wraps
from typing import Any
from typing import Any, ParamSpec, TypeVar

from pyrit.exceptions import PyritException
from pyrit.models import (
Expand All @@ -18,6 +18,9 @@

logger = logging.getLogger(__name__)

_P = ParamSpec("_P")
_T = TypeVar("_T")


def _get_rate_limit_lock(target: Any) -> asyncio.Lock:
"""Return the target's pacing lock, rebuilding it when the event loop changes."""
Expand Down Expand Up @@ -60,7 +63,9 @@ def validate_top_p(top_p: float | None) -> None:
raise PyritException(message="top_p must be between 0 and 1 (inclusive).")


def limit_requests_per_minute(func: Callable[..., Any]) -> Callable[..., Any]:
def limit_requests_per_minute(
func: Callable[_P, Awaitable[_T]],
) -> Callable[_P, Coroutine[Any, Any, _T]]:
"""
Enforce a target's request rate by serializing the delay before each request.

Expand All @@ -79,7 +84,7 @@ def limit_requests_per_minute(func: Callable[..., Any]) -> Callable[..., Any]:
"""

@wraps(func)
async def set_max_rpm_async(*args: Any, **kwargs: Any) -> Any:
async def set_max_rpm_async(*args: _P.args, **kwargs: _P.kwargs) -> _T:
self = args[0]
rpm = getattr(self, "_max_requests_per_minute", None)
if rpm and rpm > 0:
Expand Down
5 changes: 5 additions & 0 deletions tests/unit/async_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,11 @@
T = TypeVar("T")


def get_defined_tasks(*tasks: asyncio.Task[T] | None) -> list[asyncio.Task[T]]:
"""Retain owned tasks even when a test fails before creating all of them."""
return [task for task in tasks if task is not None]


async def wait_for_completion_async(*, future: asyncio.Future[T], timeout: float = 30) -> T:
"""Bound a test wait without injecting cancellation into the operation under test."""
done, _ = await asyncio.wait({future}, timeout=timeout)
Expand Down
6 changes: 4 additions & 2 deletions tests/unit/backend/test_auth_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -270,10 +270,12 @@ async def test_authenticate_request_caches_successful_authorization() -> None:

def test_auth_cache_expires_and_evicts_oldest_entry() -> None:
middleware = _make_middleware()
middleware._AUTH_CACHE_MAX_ENTRIES = 2
user = AuthenticatedUser(oid="user-1", name="Test User", email="test@example.com", groups=["allowed-group"])

with patch("pyrit.backend.middleware.auth.monotonic", return_value=100.0):
with (
patch.object(EntraAuthMiddleware, "_AUTH_CACHE_MAX_ENTRIES", 2),
patch("pyrit.backend.middleware.auth.monotonic", return_value=100.0),
):
middleware._cache_user(cache_key="first", user=user)
middleware._cache_user(cache_key="second", user=user)
middleware._cache_user(cache_key="third", user=user)
Expand Down
3 changes: 2 additions & 1 deletion tests/unit/backend/test_configuration_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

"""Tests for backend configuration file routes."""

from collections.abc import Iterator
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
Expand All @@ -23,7 +24,7 @@


@pytest.fixture
def client(compatibility_headers: dict[str, str]) -> TestClient:
def client(compatibility_headers: dict[str, str]) -> Iterator[TestClient]:
"""Create a test client for the FastAPI app."""
app.dependency_overrides[require_admin] = lambda: None
try:
Expand Down
3 changes: 2 additions & 1 deletion tests/unit/backend/test_initializer_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
Tests for backend initializer service and routes.
"""

from collections.abc import Iterator
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
Expand All @@ -26,7 +27,7 @@


@pytest.fixture
def client(compatibility_headers: dict[str, str]) -> TestClient:
def client(compatibility_headers: dict[str, str]) -> Iterator[TestClient]:
"""Create a test client for the FastAPI app."""
app.dependency_overrides[require_admin] = lambda: None
try:
Expand Down
4 changes: 2 additions & 2 deletions tests/unit/backend/test_message_send_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

import asyncio
import uuid
from collections.abc import AsyncGenerator, Generator, Iterator, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Generator, Iterator, Sequence
from contextlib import asynccontextmanager, contextmanager
from datetime import datetime
from pathlib import Path
Expand Down Expand Up @@ -106,7 +106,7 @@ def send_dependencies(mock_memory: MagicMock) -> Iterator[tuple[MagicMock, Async
@pytest.fixture
async def real_send_context(
*, sqlite_instance: SQLiteMemory, patch_central_database: MagicMock
) -> Iterator[tuple[MessageSendService, AttackResult, MockPromptTarget, Base64Converter]]:
) -> AsyncIterator[tuple[MessageSendService, AttackResult, MockPromptTarget, Base64Converter]]:
target = MockPromptTarget()
converter = Base64Converter()
ar = AttackResult(
Expand Down
3 changes: 2 additions & 1 deletion tests/unit/backend/test_scenario_progress_read_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@
TechniqueBundle,
)
from pyrit.score import Scorer, SubStringScorer
from unit.async_utils import get_defined_tasks
from unit.mocks import MockPromptTarget, get_mock_target_identifier, make_scenario_result


Expand Down Expand Up @@ -237,7 +238,7 @@ async def wait_for_snapshot_async() -> ScenarioProgressSnapshot:
assert not read_model._cache_lock.locked()
finally:
release.set()
tasks = [task for task in (owner, waiter, successor) if task is not None]
tasks = get_defined_tasks(owner, waiter, successor)
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
Expand Down
Loading
Loading