test(batches): use immutable expectations with explicit test doubles

This commit is contained in:
Yuneng Jiang 2026-09-07 16:36:27 -07:00
parent 7bff9bf9a2
commit 4ab5719ff9
No known key found for this signature in database
2 changed files with 93 additions and 96 deletions

View file

@ -31,7 +31,7 @@ FileCreateProvider = Literal[
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
]
FileDeleteProvider = Literal["openai", "azure", "gemini", "litellm_proxy", "manus", "anthropic"]
FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
import litellm
from litellm import get_secret_str

View file

@ -1,7 +1,7 @@
from builtins import ExceptionGroup
from collections.abc import Iterator
from dataclasses import dataclass, field
from collections.abc import Callable
from typing import Final
from unittest.mock import Mock, call
import pytest
from batch_cleanup import BATCH_CANCEL_TIMEOUT_SECONDS, CLEANUP_DELAYS, cleanup_batch, cleanup_file, cleanup_result
@ -15,35 +15,43 @@ MANAGED_FILE_ID: Final = "bGl0ZWxsbV9wcm94eTtmaWxlLTE="
MANAGED_BATCH_ID: Final = "bGl0ZWxsbV9wcm94eTtiYXRjaC0x"
@dataclass(frozen=True, slots=True)
class ExpectedCalls[T]:
values: Iterator[T]
def __init__(self, values: tuple[T, ...]) -> None:
self.values: Final = values
self.recorder: Final = Mock()
def __call__(self, value: T) -> None:
assert next(self.values, None) == value
self.recorder(value)
def assert_done(self) -> None:
assert tuple(self.values) == ()
assert tuple(self.recorder.call_args_list) == tuple(call(value) for value in self.values)
@dataclass(frozen=True, slots=True)
class CleanupClient:
calls: ExpectedCalls[str]
files: Iterator[Result[FileDeleteResponse]] = field(default_factory=lambda: iter(()))
batches: Iterator[Result[BatchObject]] = field(default_factory=lambda: iter(()))
cancellations: Iterator[Result[BatchObject]] = field(default_factory=lambda: iter(()))
def __init__(
self,
*,
calls: ExpectedCalls[str],
files: tuple[Result[FileDeleteResponse], ...] = (),
batches: tuple[Result[BatchObject], ...] = (),
cancellations: tuple[Result[BatchObject], ...] = (),
) -> None:
self.calls: Final = calls
self.file_response: Final[Callable[[], Result[FileDeleteResponse]]] = Mock(side_effect=files)
self.batch_response: Final[Callable[[], Result[BatchObject]]] = Mock(side_effect=batches)
self.cancel_response: Final[Callable[[], Result[BatchObject]]] = Mock(side_effect=cancellations)
def delete_file(self, file_id: str, *, key: str, provider: str | None = None) -> Result[FileDeleteResponse]:
self.calls(f"delete {provider} {file_id}")
return next(self.files)
return self.file_response()
def retrieve_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]:
self.calls(f"retrieve {provider} {batch_id}")
return next(self.batches)
return self.batch_response()
def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]:
self.calls(f"cancel {provider} {batch_id}")
return next(self.cancellations)
return self.cancel_response()
def generate_key(self, body: KeyGenerateBody) -> str:
return "test-key"
@ -68,17 +76,15 @@ class TestFileCleanup:
response: Final = Success(
status_code=200, data=FileDeleteResponse.model_validate({"id": MANAGED_FILE_ID, "object": "file"})
)
client: Final = CleanupClient(
calls=ExpectedCalls(iter((f"delete None {MANAGED_FILE_ID}",))), files=iter((response,))
)
client: Final = CleanupClient(calls=ExpectedCalls((f"delete None {MANAGED_FILE_ID}",)), files=(response,))
cleanup_file(client, MANAGED_FILE_ID, key="test-key")
client.calls.assert_done()
@pytest.mark.parametrize("file_id", ["file-1", MANAGED_FILE_ID])
def test_a_success_status_without_a_deletion_confirmation_is_rejected(self, file_id: str) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter((f"delete None {file_id}",))),
files=iter((Success(status_code=200, data=FileDeleteResponse(id=file_id)),)),
calls=ExpectedCalls((f"delete None {file_id}",)),
files=(Success(status_code=200, data=FileDeleteResponse(id=file_id)),),
)
with pytest.raises(AssertionError, match="did not confirm deletion"):
cleanup_file(client, file_id, key="test-key")
@ -88,15 +94,15 @@ class TestFileCleanup:
def test_deletes_raw_files_through_the_upload_provider(self, cap: Capability) -> None:
expected_provider: Final = cap.provider if cap.scenario in {"model_param", "provider_fallback"} else None
client: Final = CleanupClient(
calls=ExpectedCalls(iter((f"delete {expected_provider} file-1",))), files=iter((deleted_file(),))
calls=ExpectedCalls((f"delete {expected_provider} file-1",)), files=(deleted_file(),)
)
cleanup_file(client, "file-1", key="test-key", provider=cap.file_provider)
client.calls.assert_done()
def test_failed_delete_is_reported_after_remaining_resources_are_cleaned(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("delete azure file-1", "delete key test-key"))),
files=iter((UnknownApiError(status_code=403, body="secret response"),)),
calls=ExpectedCalls(("delete azure file-1", "delete key test-key")),
files=(UnknownApiError(status_code=403, body="secret response"),),
)
manager: Final = ResourceManager(client=client, strict_cleanup=True)
key: Final = manager.key()
@ -109,7 +115,7 @@ class TestFileCleanup:
def test_success_response_must_confirm_deletion(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("delete None file-1",))), files=iter((deleted_file(deleted=False),))
calls=ExpectedCalls(("delete None file-1",)), files=(deleted_file(deleted=False),)
)
with pytest.raises(AssertionError, match="did not confirm deletion"):
cleanup_file(client, "file-1", key="test-key")
@ -117,16 +123,16 @@ class TestFileCleanup:
def test_cleanup_is_idempotent_when_file_is_already_deleted(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("delete azure file-1",))),
files=iter((UnknownApiError(status_code=404, body="missing"),)),
calls=ExpectedCalls(("delete azure file-1",)),
files=(UnknownApiError(status_code=404, body="missing"),),
)
cleanup_file(client, "file-1", key="test-key", provider="azure")
client.calls.assert_done()
def test_default_resource_cleanup_keeps_existing_best_effort_behavior(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("delete None file-1", "delete key test-key"))),
files=iter((UnknownApiError(status_code=403, body="forbidden"),)),
calls=ExpectedCalls(("delete None file-1", "delete key test-key")),
files=(UnknownApiError(status_code=403, body="forbidden"),),
)
manager: Final = ResourceManager(client=client)
key: Final = manager.key()
@ -141,37 +147,39 @@ class TestCleanupRetries:
[NetworkError(message="offline"), RateLimitedError(), UnknownApiError(status_code=503, body="unavailable")],
)
def test_transient_error_retries_and_returns_success(self, failure: Result[FileDeleteResponse]) -> None:
outcomes: Final = iter((failure, deleted_file()))
delays: Final = ExpectedCalls(iter((1.0,)))
result: Final[Result[FileDeleteResponse]] = cleanup_result(lambda: next(outcomes), wait=delays)
responses: Final = (failure, deleted_file())
outcomes: Final = Mock(side_effect=responses)
delays: Final = ExpectedCalls((1.0,))
result: Final[Result[FileDeleteResponse]] = cleanup_result(outcomes, wait=delays)
assert isinstance(result, Success) and result.data.deleted
delays.assert_done()
def test_persistent_error_has_bounded_retries(self) -> None:
failure: Final = UnknownApiError(status_code=503, body="unavailable")
outcomes: Final[Iterator[Result[FileDeleteResponse]]] = iter((failure,) * (len(CLEANUP_DELAYS) + 1))
delays: Final = ExpectedCalls(iter(CLEANUP_DELAYS))
result: Final[Result[FileDeleteResponse]] = cleanup_result(lambda: next(outcomes), wait=delays)
outcomes: Final = Mock(return_value=failure)
delays: Final = ExpectedCalls(CLEANUP_DELAYS)
result: Final[Result[FileDeleteResponse]] = cleanup_result(outcomes, wait=delays)
assert result is failure
delays.assert_done()
assert next(outcomes, None) is None
assert outcomes.call_count == len(CLEANUP_DELAYS) + 1
def test_permanent_error_is_not_retried(self) -> None:
failure: Final = UnknownApiError(status_code=403, body="forbidden")
outcomes: Final = iter((failure, deleted_file()))
delays: Final = ExpectedCalls[float](iter(()))
assert cleanup_result(lambda: next(outcomes), wait=delays) is failure
responses: Final = (failure, deleted_file())
outcomes: Final = Mock(side_effect=responses)
delays: Final = ExpectedCalls[float](())
assert cleanup_result(outcomes, wait=delays) is failure
delays.assert_done()
assert isinstance(next(outcomes), Success)
assert outcomes.call_count == 1
class TestBatchCancellation:
def test_cancelling_batch_is_polled_until_terminal_without_cancelling_again(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter((f"retrieve None {MANAGED_BATCH_ID}",) * 3)),
batches=iter((batch("cancelling"), batch("cancelling"), batch("cancelled"))),
calls=ExpectedCalls((f"retrieve None {MANAGED_BATCH_ID}",) * 3),
batches=(batch("cancelling"), batch("cancelling"), batch("cancelled")),
)
delays: Final = ExpectedCalls(iter((10.0,)))
delays: Final = ExpectedCalls((10.0,))
cleanup_batch(client, MANAGED_BATCH_ID, key="test-key", wait=delays)
client.calls.assert_done()
delays.assert_done()
@ -179,23 +187,22 @@ class TestBatchCancellation:
def test_cancellation_timeout_is_reported_but_file_and_key_cleanup_still_run(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(
iter(
(
f"retrieve None {MANAGED_BATCH_ID}",
f"retrieve None {MANAGED_BATCH_ID}",
"delete None file-1",
"delete key test-key",
)
(
f"retrieve None {MANAGED_BATCH_ID}",
f"retrieve None {MANAGED_BATCH_ID}",
"delete None file-1",
"delete key test-key",
)
),
batches=iter((batch("cancelling"), batch("cancelling"))),
files=iter((deleted_file(),)),
batches=(batch("cancelling"), batch("cancelling")),
files=(deleted_file(),),
)
ticks: Final = iter((0.0, BATCH_CANCEL_TIMEOUT_SECONDS))
times: Final = (0.0, BATCH_CANCEL_TIMEOUT_SECONDS)
ticks: Final[Callable[[], float]] = Mock(side_effect=times)
manager: Final = ResourceManager(client=client, strict_cleanup=True)
key: Final = manager.key()
manager.defer(lambda: cleanup_file(client, "file-1", key=key))
manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=lambda: next(ticks)))
manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=ticks))
with pytest.raises(ExceptionGroup) as caught:
manager.teardown()
assert "cancellation did not finish" in str(caught.value.exceptions[0])
@ -203,17 +210,15 @@ class TestBatchCancellation:
@pytest.mark.parametrize("status", ["completed", "failed", "expired", "cancelled"])
def test_inactive_batch_needs_no_cancellation(self, status: str) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("retrieve None batch-1",))), batches=iter((batch(status),))
)
client: Final = CleanupClient(calls=ExpectedCalls(("retrieve None batch-1",)), batches=(batch(status),))
cleanup_batch(client, "batch-1", key="test-key")
client.calls.assert_done()
def test_active_batch_is_cancelled_through_its_provider(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("retrieve azure batch-1", "cancel azure batch-1"))),
batches=iter((batch("in_progress"), batch("cancelled"))),
cancellations=iter((batch("cancelling"),)),
calls=ExpectedCalls(("retrieve azure batch-1", "cancel azure batch-1")),
batches=(batch("in_progress"), batch("cancelled")),
cancellations=(batch("cancelling"),),
)
cleanup_batch(client, "batch-1", key="test-key", provider="azure")
client.calls.assert_done()
@ -225,23 +230,21 @@ class TestBatchCancellation:
) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(
iter(
(
f"retrieve vertex_ai {batch_id}",
f"cancel vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
"delete vertex_ai file-1",
"delete key test-key",
)
(
f"retrieve vertex_ai {batch_id}",
f"cancel vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
f"retrieve vertex_ai {batch_id}",
"delete vertex_ai file-1",
"delete key test-key",
)
),
batches=iter((batch("validating"), batch(pending_status), batch(pending_status), batch("cancelled"))),
cancellations=iter((batch(pending_status),)),
files=iter((deleted_file(),)),
batches=(batch("validating"), batch(pending_status), batch(pending_status), batch("cancelled")),
cancellations=(batch(pending_status),),
files=(deleted_file(),),
)
delays: Final = ExpectedCalls(iter((10.0, 10.0)))
delays: Final = ExpectedCalls((10.0, 10.0))
manager: Final = ResourceManager(client=client, strict_cleanup=True)
key: Final = manager.key()
manager.defer(lambda: cleanup_file(client, "file-1", key=key, provider="vertex_ai"))
@ -255,28 +258,22 @@ class TestBatchCancellation:
self, output_delete_fails: bool
) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(
iter(("retrieve openai batch-1", "delete openai file-output", "delete openai file-error"))
),
batches=iter(
(
Success(
status_code=200,
data=BatchObject(
id="batch-1",
status="completed",
input_file_id="file-input",
output_file_id="file-output",
error_file_id="file-error",
),
calls=ExpectedCalls(("retrieve openai batch-1", "delete openai file-output", "delete openai file-error")),
batches=(
Success(
status_code=200,
data=BatchObject(
id="batch-1",
status="completed",
input_file_id="file-input",
output_file_id="file-output",
error_file_id="file-error",
),
)
),
),
files=iter(
(
UnknownApiError(status_code=403, body="forbidden") if output_delete_fails else deleted_file(),
deleted_file(),
)
files=(
UnknownApiError(status_code=403, body="forbidden") if output_delete_fails else deleted_file(),
deleted_file(),
),
)
if output_delete_fails:
@ -289,9 +286,9 @@ class TestBatchCancellation:
@pytest.mark.parametrize("status", ["completed", "in_progress"])
def test_cancellation_conflict_is_accepted_only_when_batch_became_inactive(self, status: str) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(iter(("retrieve None batch-1", "cancel None batch-1", "retrieve None batch-1"))),
batches=iter((batch("in_progress"), batch(status))),
cancellations=iter((UnknownApiError(status_code=409, body="conflict"),)),
calls=ExpectedCalls(("retrieve None batch-1", "cancel None batch-1", "retrieve None batch-1")),
batches=(batch("in_progress"), batch(status)),
cancellations=(UnknownApiError(status_code=409, body="conflict"),),
)
if status == "completed":
cleanup_batch(client, "batch-1", key="test-key")