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
15 changes: 15 additions & 0 deletions doc/contributing/5_unit_tests.md
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,21 @@ Testing is an art to get right! But here are some best practices in terms of uni

Not all of our current tests follow these practices (we're working on it!) But for some good examples, see [test_tts_send_prompt_file_save_async](../../tests/unit/prompt_target/target/test_tts_target.py), which has many of these best practices incorporated in the test.

## Async timing and cancellation

Use events to coordinate concurrent operations and assert their ordering or concurrency bounds.
Timeouts that only prevent a test from hanging should allow for busy CI runners, rather than
acting as performance assertions.

When observing an operation's cancellation or cleanup, use `wait_for_completion_async` from
`unit.async_utils`. Unlike `asyncio.wait_for`, its watchdog does not send another cancellation
request to the operation when the wait expires. Release blocked workers and drain owned tasks
in `finally` so a failed assertion does not leave background work behind.

For deadline tests, expire a real `asyncio.Timeout` with `reschedule` once the operation reaches
the intended pending await. Check the configured timeout arguments, cancellation, cleanup, and
original outcome. This avoids short wall-clock deadlines expiring during unrelated setup.

## SQLite memory fixtures

`sqlite_instance` stays function-scoped. Each test gets a fresh in-memory database and results directory, and its SQLite singleton and CentralMemory registrations are restored afterward. The fixture owns disposal of its memory instance instead of registering process-exit cleanup callbacks for every test.
Expand Down
2 changes: 1 addition & 1 deletion frontend/e2e/numeric-controls.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -79,8 +79,8 @@ async function clickNativeSpinner(input: Locator, direction: 'up' | 'down'): Pro
if (!box) throw new Error('Expected a visible numeric input.')
const paddingRight = await input.evaluate((element) => parseFloat(getComputedStyle(element).paddingRight))
// Chromium's native spinner is a UA shadow control, not an accessible button.
// Native spinners auto-repeat when held, so a single-step check must not hold the mouse down.
await input.click({
delay: 200,
position: { x: box.width - paddingRight - 8, y: box.height / 2 + (direction === 'up' ? -4 : 4) },
})
}
Expand Down
15 changes: 15 additions & 0 deletions tests/unit/async_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import asyncio
from typing import TypeVar

T = TypeVar("T")


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)
if not done:
raise TimeoutError("The operation under test did not complete before the test watchdog expired.")
return future.result()
6 changes: 3 additions & 3 deletions tests/unit/backend/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ async def test_health_responds_while_database_operation_is_pending(sqlite_instan

def wait_in_database() -> int:
started.set()
if not release.wait(timeout=10):
if not release.wait(timeout=60):
raise RuntimeError("Database wait was not released")
return 1

Expand All @@ -55,9 +55,9 @@ def wait_in_database() -> int:
)
query = asyncio.create_task(session.execute(text("SELECT wait_in_database()")))
try:
assert await asyncio.to_thread(started.wait, 5)
assert await asyncio.to_thread(started.wait, 30)
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
response = await asyncio.wait_for(client.get("/api/health"), timeout=2)
response = await asyncio.wait_for(client.get("/api/health"), timeout=30)
assert response.status_code == 200
assert response.json()["status"] == "healthy"
assert not query.done()
Expand Down
38 changes: 24 additions & 14 deletions tests/unit/backend/test_scenario_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -769,16 +769,21 @@ async def estimate_async(
service._run_default_estimate_async = AsyncMock(side_effect=estimate_async)

catalog_task = asyncio.create_task(service.list_scenarios_async())
await asyncio.wait_for(two_started.wait(), timeout=2)
await asyncio.sleep(0)
try:
await asyncio.wait_for(two_started.wait(), timeout=30)
await asyncio.sleep(0)

assert service._run_default_estimate_async.await_count == 2
release.set()
result = await catalog_task
assert service._run_default_estimate_async.await_count == 2
release.set()
result = await catalog_task

assert maximum_active == 2
assert service._run_default_estimate_async.await_count == 3
assert all(item.default_run_size == estimate for item in result.items)
assert maximum_active == 2
assert service._run_default_estimate_async.await_count == 3
assert all(item.default_run_size == estimate for item in result.items)
finally:
release.set()
await service.close_async()
await asyncio.gather(catalog_task, return_exceptions=True)

async def test_catalog_queue_wait_does_not_start_execution_timeout(self) -> None:
"""A queued catalog estimate starts its timeout only after acquiring capacity."""
Expand Down Expand Up @@ -1259,13 +1264,18 @@ async def estimate_async(
)
for index in range(3)
]
await asyncio.wait_for(two_started.wait(), timeout=1)
await asyncio.sleep(0)
try:
await asyncio.wait_for(two_started.wait(), timeout=30)
await asyncio.sleep(0)

assert service._estimate_configured_run_size_async.await_count == 2
release.set()
assert await asyncio.gather(*tasks) == [estimate, estimate, estimate]
assert maximum_active == 2
assert service._estimate_configured_run_size_async.await_count == 2
release.set()
assert await asyncio.gather(*tasks) == [estimate, estimate, estimate]
assert maximum_active == 2
finally:
release.set()
await service.close_async()
await asyncio.gather(*tasks, return_exceptions=True)

async def test_metadata_catalog_remains_responsive_during_estimate(self) -> None:
"""Metadata-only catalog requests do not wait for running estimates."""
Expand Down
65 changes: 47 additions & 18 deletions tests/unit/executor/promptgen/test_target_objective_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,21 @@ def _isolate_pre_send_io(generator: TargetObjectiveGenerator) -> Iterator[AsyncM
yield setup


@contextmanager
def _controlled_timeouts() -> Iterator[list[tuple[float | None, asyncio.Timeout]]]:
"""Expire real asyncio deadlines at the intended await, not during unrelated setup."""
timeouts: list[tuple[float | None, asyncio.Timeout]] = []
original_timeout = asyncio.timeout

def capture_timeout(delay: float | None) -> asyncio.Timeout:
timeout = original_timeout(None)
timeouts.append((delay, timeout))
return timeout

with patch("pyrit.executor.promptgen.target_objective_generator.asyncio.timeout", new=capture_timeout):
yield timeouts


@pytest.mark.usefixtures("patch_central_database")
class TestTargetObjectiveGenerator:
async def test_valid_batch_and_evidence_async(self, sqlite_instance: MemoryInterface) -> None:
Expand Down Expand Up @@ -231,41 +246,47 @@ def test_nontext_system_prompt_rejected(self) -> None:
target=MockPromptTarget(), system_prompt=SeedPrompt(value="image.png", data_type="image_path")
)

@pytest.mark.timeout(30)
async def test_timeout_bounds_pending_send_and_cleanup_async(self, caplog: pytest.LogCaptureFixture) -> None:
generator = TargetObjectiveGenerator(target=MockPromptTarget(), timeout_seconds=1)
cancelled = asyncio.Event()
cleanup_cancelled = asyncio.Event()

async def wait_forever_async(**kwargs: object) -> None:
try:
timeouts[0][1].reschedule(asyncio.get_running_loop().time())
await asyncio.Event().wait()
finally:
cancelled.set()

async def reset_async(*, conversation_id: str) -> None:
try:
timeouts[-1][1].reschedule(asyncio.get_running_loop().time())
await asyncio.Event().wait()
finally:
cleanup_cancelled.set()

with (
_isolate_pre_send_io(generator) as setup,
_controlled_timeouts() as timeouts,
patch.object(generator, "_CLEANUP_TIMEOUT_SECONDS", 0.01),
patch.object(
generator._normalizer, "send_prompt_async", new_callable=AsyncMock, side_effect=wait_forever_async
) as send,
patch.object(generator._target, "reset_conversation_async", side_effect=reset_async) as reset,
):
async with asyncio.timeout(3):
with pytest.raises(TimeoutError):
await generator.execute_async(instructions="Test", count=2)
with pytest.raises(TimeoutError):
await generator.execute_async(instructions="Test", count=2)
assert [delay for delay, _ in timeouts] == [1, 0.01]
assert all(timeout.expired() for _, timeout in timeouts)
send.assert_awaited_once()
assert cancelled.is_set()
assert cleanup_cancelled.is_set()
setup.assert_awaited_once()
reset.assert_awaited_once()
assert "Timed out resetting generation conversation" in caplog.text

@pytest.mark.timeout(30)
@pytest.mark.parametrize(
("pending_method", "conversation_initialized"),
[
Expand All @@ -284,21 +305,24 @@ async def test_timeout_before_send_preserves_cleanup_contract_async(

async def wait_forever_async(**kwargs: object) -> None:
try:
timeouts[0][1].reschedule(asyncio.get_running_loop().time())
await asyncio.Event().wait()
finally:
cancelled.set()

with (
_isolate_pre_send_io(generator),
_controlled_timeouts() as timeouts,
patch.object(
pending_owner, pending_method, new_callable=AsyncMock, side_effect=wait_forever_async
) as pending,
patch.object(generator._normalizer, "send_prompt_async", new_callable=AsyncMock) as send,
patch.object(generator._target, "reset_conversation_async", new_callable=AsyncMock) as reset,
):
async with asyncio.timeout(3):
with pytest.raises(TimeoutError):
await generator.execute_with_context_async(context=context)
with pytest.raises(TimeoutError):
await generator.execute_with_context_async(context=context)
assert timeouts[0][0] == 1
assert timeouts[0][1].expired()
pending.assert_awaited_once()
assert cancelled.is_set()
assert context._used
Expand All @@ -309,6 +333,7 @@ async def wait_forever_async(**kwargs: object) -> None:
else:
reset.assert_not_awaited()

@pytest.mark.timeout(30)
@pytest.mark.parametrize("failure", [None, ConnectionError("Generation failed"), asyncio.CancelledError()])
async def test_cleanup_timeout_preserves_outcome_async(
self, *, failure: BaseException | None, caplog: pytest.LogCaptureFixture
Expand All @@ -319,12 +344,14 @@ async def test_cleanup_timeout_preserves_outcome_async(

async def reset_async(*, conversation_id: str) -> None:
try:
timeouts[-1][1].reschedule(asyncio.get_running_loop().time())
await asyncio.Event().wait()
finally:
cleanup_cancelled.set()

with (
_isolate_pre_send_io(generator) as setup,
_controlled_timeouts() as timeouts,
patch.object(generator, "_CLEANUP_TIMEOUT_SECONDS", 0.01),
patch.object(
generator._normalizer,
Expand All @@ -335,18 +362,20 @@ async def reset_async(*, conversation_id: str) -> None:
) as send,
patch.object(generator._target, "reset_conversation_async", side_effect=reset_async) as reset,
):
async with asyncio.timeout(1):
if failure is None:
result = await generator.execute_async(instructions="Test", count=1)
assert result.objectives == ["A goal"]
elif isinstance(failure, asyncio.CancelledError):
with pytest.raises(asyncio.CancelledError) as error:
await generator.execute_async(instructions="Test", count=1)
assert error.value is failure
else:
with pytest.raises(RuntimeError) as generation_error:
await generator.execute_async(instructions="Test", count=1)
assert generation_error.value.__cause__ is failure
if failure is None:
result = await generator.execute_async(instructions="Test", count=1)
assert result.objectives == ["A goal"]
elif isinstance(failure, asyncio.CancelledError):
with pytest.raises(asyncio.CancelledError) as error:
await generator.execute_async(instructions="Test", count=1)
assert error.value is failure
else:
with pytest.raises(RuntimeError) as generation_error:
await generator.execute_async(instructions="Test", count=1)
assert generation_error.value.__cause__ is failure
assert [delay for delay, _ in timeouts] == [generator._timeout_seconds, 0.01]
assert not timeouts[0][1].expired()
assert timeouts[1][1].expired()
reset.assert_awaited_once()
assert cleanup_cancelled.is_set()
setup.assert_awaited_once()
Expand Down
Loading
Loading