Merge pull request #40161 from BerriAI/litellm_batch_e2e_cleanup

fix(batches): clean up E2E resources across providers
This commit is contained in:
yuneng-jiang 2026-09-08 22:53:48 -07:00 committed by GitHub
commit 802e526cf9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 949 additions and 108 deletions

View file

@ -1801,7 +1801,16 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Remove conflicting keys from data to avoid duplicate keyword arguments
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
for model_id, model_file_id in specific_model_file_id_mapping.items():
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
delete_data = {
**{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
**(
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
if credentials is not None
else {}
),
}
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
stored_file_object = await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
@ -1812,7 +1821,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
prom_logger.record_managed_file_deleted(result="success")
if stored_file_object:
return stored_file_object
return OpenAIFileObject.model_validate(stored_file_object).model_copy(update={"id": file_id})
elif delete_response:
delete_response.id = file_id
return delete_response

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

@ -7,13 +7,13 @@ from contextlib import suppress
from functools import cache
from itertools import chain
from types import MappingProxyType
from typing import Any, Final, TypeAlias, TypedDict
from typing import Any, Final, Literal, TypeAlias, TypedDict
from urllib.parse import unquote
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
from pydantic import BaseModel, ConfigDict, TypeAdapter
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
@ -60,11 +60,12 @@ from litellm.utils import get_llm_provider
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resolve_s3_encryption_key_id
# litellm_params key used to hand the SigV4-signed GET headers from
# `transform_file_content_request` to `validate_environment` (the only hook
# the shared file-content HTTP handler exposes for setting request headers).
# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
S3_SIGNED_GET_HEADERS_PARAM: Final = "_s3_signed_get_headers"
S3_SIGNED_REQUEST_HEADERS_PARAM: Final = "_s3_signed_request_headers"
class _S3DeleteContext(BaseModel):
file_id: str = Field(min_length=1)
# litellm_params key carrying the size of the body uploaded to S3, handed from
# `transform_create_file_request` to `transform_create_file_response`.
@ -291,7 +292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
) -> dict:
result: Final[dict[str, object]] = {}
result.update(headers)
signed_headers: Final = litellm_params.pop(S3_SIGNED_GET_HEADERS_PARAM, None)
signed_headers: Final = litellm_params.pop(S3_SIGNED_REQUEST_HEADERS_PARAM, None)
if isinstance(signed_headers, Mapping):
result.update(signed_headers) # any-ok: untyped handoff headers
# otherwise no extra headers - AWS credentials are handled by BaseAWSLLM
@ -1187,18 +1188,27 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
def transform_delete_file_request(
self,
file_id: str,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
raise NotImplementedError("BedrockFilesConfig does not support file deletion")
optional_params: Mapping[str, object],
litellm_params: MutableMapping[str, object],
) -> tuple[str, dict[str, str]]:
return self._transform_s3_file_request(
file_id=file_id, method="DELETE", optional_params=optional_params, litellm_params=litellm_params
)
def transform_delete_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
litellm_params: Mapping[str, object],
) -> FileDeleted:
raise NotImplementedError("BedrockFilesConfig does not support file deletion")
if raw_response.status_code != 204:
raise BedrockError(
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
message=raw_response.text or f"S3 file deletion returned HTTP {raw_response.status_code}",
headers=raw_response.headers,
)
context: Final = _S3DeleteContext.model_validate(logging_obj.model_call_details.get("additional_args"))
return FileDeleted(id=context.file_id, deleted=True, object="file")
def transform_list_files_request(
self,
@ -1233,6 +1243,18 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if not file_id:
raise ValueError("file_id is required for Bedrock file content retrieval")
return self._transform_s3_file_request(
file_id=file_id, method="GET", optional_params=optional_params, litellm_params=litellm_params
)
def _transform_s3_file_request(
self,
*,
file_id: str,
method: Literal["GET", "DELETE"],
optional_params: Mapping[str, object],
litellm_params: MutableMapping[str, object],
) -> tuple[str, dict[str, str]]:
s3_uri: Final = extract_s3_uri_from_file_id(file_id)
bucket_name, object_key = _validate_file_id_against_configured_buckets(
s3_uri=s3_uri,
@ -1240,40 +1262,32 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(litellm_params),
)
# The shared file-content handler passes optional_params={}, so AWS
# credentials/region arrive via litellm_params here (unlike the upload
# path). s3_region_name wins over aws_region_name, same priority as
# get_complete_file_url above.
merged_params: Final[dict[str, object]] = {}
merged_params.update(litellm_params)
merged_params.update(optional_params)
request_params: Final = _BedrockS3RequestParams.model_validate(merged_params)
request_params: Final = _BedrockS3RequestParams.model_validate({**litellm_params, **optional_params})
region_preference: Final = request_params.s3_region_name or request_params.aws_region_name
region_params: Final[dict[str, str | None]] = {"aws_region_name": region_preference}
aws_region_name: Final = self._get_aws_region_name(optional_params=region_params, model="")
s3_endpoint_url = (
s3_endpoint_url: Final = (
request_params.s3_endpoint_url or f"https://s3.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
).rstrip("/")
url: Final = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM] = self._sign_s3_request_without_body(
api_base=url,
aws_region_name=aws_region_name,
request_params=request_params,
method=method,
)
return url, {}
def _sign_s3_get_request(
def _sign_s3_request_without_body(
self,
api_base: str,
aws_region_name: str,
request_params: _BedrockS3RequestParams,
method: Literal["GET", "DELETE"] = "GET",
) -> dict[str, str]:
"""
SigV4-sign an S3 GetObject request, mirroring `_sign_s3_request` (PUT).
"""
try:
import hashlib
@ -1297,7 +1311,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()
aws_request: Final = AWSRequest( # any-ok: botocore AWSRequest is untyped
method="GET",
method=method,
url=api_base,
headers={"x-amz-content-sha256": empty_body_hash},
)

View file

@ -120,6 +120,36 @@ create traverse gateway -> gateway -> OpenAI (LIT-5347, PR #36240). The pin:
nested managed ids round-trip retrieve. This self-chaining only needs the proxy to
reach its own `PROXY_BASE_URL`, which holds both locally and on the e2e stage.
## Cleanup
Batch teardown cancels active batches before deleting their input files and keys.
Raw file IDs from both `model_param` and `provider_fallback` uploads use the upload
provider when deleted. Model-encoded and managed file IDs route themselves
File deletion and batch cancellation check their responses and retry transient
failures up to three times. Teardown attempts every registered cleanup before
reporting failures as test errors. Already deleted files and batches that are
terminal are safe to clean up again. Managed batch cancellation polls for up to eleven minutes
before input deletion: the ten-minute provider window plus a propagation margin.
Accepted cancellation may still report validating or in_progress while the provider
updates its state. Raw and model-encoded batches are polled until cancelling or
terminal before input deletion. OpenAI and Azure lifecycle cleanup also deletes
output and error files returned by terminal batches. Bedrock deletion uses a signed S3 DELETE
restricted to the configured storage buckets and managed file prefixes. The low-RPM
test submits with its restricted key and cleans up with the test administrator key
Managed deletion forwards the deployment's trusted bucket configuration and returns
the requested managed file ID even when stored output metadata carries a provider ID
Azure input uploads request `expires_after` anchored to `created_at` with
`seconds=1209600`, and the lifecycle tests check the returned expiry. This is a
fallback for interrupted runs: immediate deletion remains the normal cleanup.
Azure's minimum supported native expiry is 14 days, so a three-day expiry cannot
be requested through its Files API
The Azure entry in `files_settings` must use `api_version: 2025-04-01-preview`
for raw uploads to honor expiry, matching the batch deployment's API version
## Terminal state + cost write-back (cross-run marker baton)
The 24h completion window rules out submit-and-wait inside one run, so

View file

@ -0,0 +1,140 @@
from builtins import ExceptionGroup
from collections.abc import Callable
from itertools import count
from time import monotonic, sleep
from typing import Final, Protocol
from batch_client import BatchObject, FileDeleteResponse
from capabilities import is_managed_id
from e2e_http import NetworkError, RateLimitedError, Result, Success, UnknownApiError
from pydantic import BaseModel
CLEANUP_DELAYS: Final = (1.0, 2.0, 4.0)
BATCH_TERMINAL_STATUSES: Final = frozenset({"completed", "failed", "expired", "cancelled"})
BATCH_PENDING_STATUSES: Final = frozenset({"validating", "in_progress", "finalizing", "cancelling"})
BATCH_CANCEL_TIMEOUT_SECONDS: Final = 660.0
BATCH_CANCEL_POLL_SECONDS: Final = 10.0
class BatchCleanupClient(Protocol):
def delete_file(self, file_id: str, *, key: str, provider: str | None = None) -> Result[FileDeleteResponse]: ...
def retrieve_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ...
def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ...
def cleanup_result[R: BaseModel](
action: Callable[[], Result[R]], *, wait: Callable[[float], None] = sleep
) -> Result[R]:
for delay, result in ((delay, action()) for delay in CLEANUP_DELAYS):
match result:
case NetworkError() | RateLimitedError():
wait(delay)
case UnknownApiError(status_code=code) if code in {408, 429, 500, 502, 503, 504}:
wait(delay)
case _:
return result
return action()
def _require_cleanup_success[R: BaseModel](result: Result[R], operation: str) -> R:
match result:
case Success(data=data):
return data
case UnknownApiError(status_code=code):
raise AssertionError(f"{operation} failed: HTTP {code}")
case _:
raise AssertionError(f"{operation} failed: {result.kind}")
def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider: str | None = None) -> None:
result: Final = cleanup_result(lambda: client.delete_file(file_id, key=key, provider=provider))
if isinstance(result, UnknownApiError) and result.status_code == 404:
return
deleted: Final = _require_cleanup_success(result, f"Delete file {file_id}")
assert deleted.deleted is True or (
deleted.deleted is None and is_managed_id(file_id) and deleted.id == file_id and deleted.object == "file"
), f"Delete file {file_id} did not confirm deletion"
def cleanup_batch(
client: BatchCleanupClient,
batch_id: str,
*,
key: str,
provider: str | None = None,
delete_output_files: bool = False,
wait: Callable[[float], None] = sleep,
clock: Callable[[], float] = monotonic,
) -> None:
needs_terminal_state: Final = is_managed_id(batch_id)
fetched: Final = _require_cleanup_success(
cleanup_result(lambda: client.retrieve_batch(batch_id, key=key, provider=provider)),
f"Retrieve batch {batch_id} for cleanup",
)
if fetched.status in BATCH_TERMINAL_STATUSES:
if delete_output_files:
_cleanup_batch_outputs(client, fetched, key=key, provider=provider)
return
if fetched.status == "cancelling" and not needs_terminal_state:
return
result: Final = (
Success(status_code=200, data=fetched)
if fetched.status == "cancelling"
else cleanup_result(lambda: client.cancel_batch(batch_id, key=key, provider=provider))
)
conflicted: Final = isinstance(result, UnknownApiError) and result.status_code in {400, 409}
if not conflicted:
cancelled: Final = _require_cleanup_success(result, f"Cancel batch {batch_id}")
assert cancelled.status in BATCH_TERMINAL_STATUSES | BATCH_PENDING_STATUSES, (
f"Cancel batch {batch_id} left status {cancelled.status}"
)
if cancelled.status in BATCH_TERMINAL_STATUSES:
if delete_output_files:
_cleanup_batch_outputs(client, cancelled, key=key, provider=provider)
return
if cancelled.status == "cancelling" and not needs_terminal_state:
return
deadline: Final = clock() + BATCH_CANCEL_TIMEOUT_SECONDS
for current in (
_require_cleanup_success(
cleanup_result(lambda: client.retrieve_batch(batch_id, key=key, provider=provider)),
f"Retrieve batch {batch_id} after cancellation",
)
for _ in count()
):
if current.status in BATCH_TERMINAL_STATUSES:
if delete_output_files:
_cleanup_batch_outputs(client, current, key=key, provider=provider)
return
assert current.status in ({"cancelling"} if conflicted else BATCH_PENDING_STATUSES), (
f"Cancel batch {batch_id} left status {current.status}"
)
if current.status == "cancelling" and not needs_terminal_state:
return
assert clock() < deadline, (
f"Batch {batch_id} cancellation did not finish within {BATCH_CANCEL_TIMEOUT_SECONDS}s"
)
wait(BATCH_CANCEL_POLL_SECONDS)
def _cleanup_batch_outputs(client: BatchCleanupClient, batch: BatchObject, *, key: str, provider: str | None) -> None:
errors: Final = tuple(
error
for file_id in dict.fromkeys((batch.output_file_id, batch.error_file_id))
if file_id is not None and file_id != batch.input_file_id
if (error := _output_cleanup_error(client, file_id, key=key, provider=provider)) is not None
)
if errors:
raise ExceptionGroup(f"Batch {batch.id} output cleanup failed", errors)
def _output_cleanup_error(
client: BatchCleanupClient, file_id: str, *, key: str, provider: str | None
) -> Exception | None:
try:
cleanup_file(client, file_id, key=key, provider=provider)
except Exception as error:
return error
return None

View file

@ -13,8 +13,9 @@ co-located here because only this suite uses them.
from __future__ import annotations
from dataclasses import dataclass
from typing import Final, Literal
from pydantic import BaseModel
from pydantic import BaseModel, Field
from proxy_client import ProxyClient
from e2e_http import (
@ -27,6 +28,18 @@ from e2e_http import (
from models import LiteLLMParamsBody
UPLOAD_FILENAME = "batch_input.jsonl"
AZURE_FILE_EXPIRY_SECONDS: Final = 14 * 24 * 60 * 60
class ExpiringFileUploadForm(FileUploadForm):
expires_after_anchor: Literal["created_at"] = Field(default="created_at", alias="expires_after[anchor]")
expires_after_seconds: int = Field(default=AZURE_FILE_EXPIRY_SECONDS, alias="expires_after[seconds]")
def batch_upload_form(provider: str, *, target_model_names: str | None = None) -> FileUploadForm:
if provider == "azure":
return ExpiringFileUploadForm(target_model_names=target_model_names)
return FileUploadForm(target_model_names=target_model_names)
class FileObject(BaseModel):
@ -37,6 +50,7 @@ class FileObject(BaseModel):
bytes: int | None = None
status: str | None = None
created_at: int | None = None
expires_at: int | None = None
class FileList(BaseModel):
@ -85,7 +99,7 @@ class BatchList(BaseModel):
class FileDeleteResponse(BaseModel):
id: str
object: str | None = None
deleted: bool
deleted: bool | None = None
class BatchCreateBody(BaseModel):

View file

@ -108,6 +108,10 @@ class Capability:
def id(self) -> str:
return f"{self.provider}-{self.scenario}"
@property
def file_provider(self) -> str | None:
return self.provider if self.scenario in {"model_param", "provider_fallback"} else None
@property
def jsonl_model(self) -> str:
# Always the provider deployment name. Unified routes via

View file

@ -13,7 +13,7 @@ the proxy config.
from __future__ import annotations
import os
from typing import Iterator
from typing import Final, Iterator
import pytest
@ -21,6 +21,7 @@ from batch_client import BatchClient, build_client
from capabilities import PROVIDERS
from e2e_config import MANAGED_FILES_OPT_IN_ENV
from e2e_http import NoBody
from lifecycle import ResourceManager
from proxy_client import ProxyClient
@ -52,6 +53,13 @@ def client(proxy: ProxyClient) -> BatchClient:
return build_client(proxy)
@pytest.fixture
def resources(client: BatchClient) -> Iterator[ResourceManager]:
manager: Final = ResourceManager(client=client.proxy, strict_cleanup=True)
yield manager
manager.teardown()
@pytest.fixture(scope="session")
def batch_deployments(client: BatchClient) -> Iterator[None]:
probe = client.proxy.probe("/health/liveliness", params=NoBody())

View file

@ -0,0 +1,313 @@
from builtins import ExceptionGroup
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
from batch_client import AZURE_FILE_EXPIRY_SECONDS, BatchObject, FileDeleteResponse, batch_upload_form
from capabilities import CAPABILITIES, Capability
from e2e_http import NetworkError, RateLimitedError, Result, Success, UnknownApiError
from lifecycle import ResourceManager
from models import KeyGenerateBody
MANAGED_FILE_ID: Final = "bGl0ZWxsbV9wcm94eTtmaWxlLTE="
MANAGED_BATCH_ID: Final = "bGl0ZWxsbV9wcm94eTtiYXRjaC0x"
class ExpectedCalls[T]:
def __init__(self, values: tuple[T, ...]) -> None:
self.values: Final = values
self.recorder: Final = Mock()
def __call__(self, value: T) -> None:
self.recorder(value)
def assert_done(self) -> None:
assert tuple(self.recorder.call_args_list) == tuple(call(value) for value in self.values)
class CleanupClient:
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 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 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 self.cancel_response()
def generate_key(self, body: KeyGenerateBody) -> str:
return "test-key"
def delete_key(self, key: str) -> None:
self.calls(f"delete key {key}")
def delete_customers(self, user_ids: list[str]) -> None:
self.calls(f"delete customers {user_ids}")
def batch(status: str) -> Success[BatchObject]:
return Success(status_code=200, data=BatchObject(id="batch-1", status=status))
def deleted_file(*, deleted: bool = True) -> Success[FileDeleteResponse]:
return Success(status_code=200, data=FileDeleteResponse(id="file-1", deleted=deleted))
class TestFileCleanup:
def test_managed_delete_accepts_the_deleted_file_object(self) -> None:
response: Final = Success(
status_code=200, data=FileDeleteResponse.model_validate({"id": MANAGED_FILE_ID, "object": "file"})
)
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((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")
client.calls.assert_done()
@pytest.mark.parametrize("cap", CAPABILITIES, ids=[cap.id for cap in CAPABILITIES])
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((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(("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()
manager.defer(lambda: cleanup_file(client, "file-1", key=key, provider="azure"))
with pytest.raises(ExceptionGroup) as caught:
manager.teardown()
client.calls.assert_done()
assert len(caught.value.exceptions) == 1
assert str(caught.value.exceptions[0]) == "Delete file file-1 failed: HTTP 403"
def test_success_response_must_confirm_deletion(self) -> None:
client: Final = CleanupClient(
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")
client.calls.assert_done()
def test_cleanup_is_idempotent_when_file_is_already_deleted(self) -> None:
client: Final = CleanupClient(
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(("delete None file-1", "delete key test-key")),
files=(UnknownApiError(status_code=403, body="forbidden"),),
)
manager: Final = ResourceManager(client=client)
key: Final = manager.key()
manager.defer(lambda: cleanup_file(client, "file-1", key=key))
manager.teardown()
client.calls.assert_done()
class TestCleanupRetries:
@pytest.mark.parametrize(
"failure",
[NetworkError(message="offline"), RateLimitedError(), UnknownApiError(status_code=503, body="unavailable")],
)
def test_transient_error_retries_and_returns_success(self, failure: Result[FileDeleteResponse]) -> None:
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 = 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 outcomes.call_count == len(CLEANUP_DELAYS) + 1
def test_permanent_error_is_not_retried(self) -> None:
failure: Final = UnknownApiError(status_code=403, body="forbidden")
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 outcomes.call_count == 1
class TestBatchCancellation:
def test_cancelling_batch_is_polled_until_terminal_without_cancelling_again(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls((f"retrieve None {MANAGED_BATCH_ID}",) * 3),
batches=(batch("cancelling"), batch("cancelling"), batch("cancelled")),
)
delays: Final = ExpectedCalls((10.0,))
cleanup_batch(client, MANAGED_BATCH_ID, key="test-key", wait=delays)
client.calls.assert_done()
delays.assert_done()
def test_cancellation_timeout_is_reported_but_file_and_key_cleanup_still_run(self) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(
(
f"retrieve None {MANAGED_BATCH_ID}",
f"retrieve None {MANAGED_BATCH_ID}",
"delete None file-1",
"delete key test-key",
)
),
batches=(batch("cancelling"), batch("cancelling")),
files=(deleted_file(),),
)
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=ticks))
with pytest.raises(ExceptionGroup) as caught:
manager.teardown()
assert "cancellation did not finish" in str(caught.value.exceptions[0])
client.calls.assert_done()
@pytest.mark.parametrize("status", ["completed", "failed", "expired", "cancelled"])
def test_inactive_batch_needs_no_cancellation(self, status: str) -> None:
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(("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()
@pytest.mark.parametrize("batch_id", ["batch-1", MANAGED_BATCH_ID])
@pytest.mark.parametrize("pending_status", ["validating", "in_progress"])
def test_accepted_cancellation_waits_through_stale_provider_status(
self, batch_id: str, pending_status: str
) -> None:
client: Final = CleanupClient(
calls=ExpectedCalls(
(
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=(batch("validating"), batch(pending_status), batch(pending_status), batch("cancelled")),
cancellations=(batch(pending_status),),
files=(deleted_file(),),
)
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"))
manager.defer(lambda: cleanup_batch(client, batch_id, key=key, provider="vertex_ai", wait=delays))
manager.teardown()
client.calls.assert_done()
delays.assert_done()
@pytest.mark.parametrize("output_delete_fails", [False, True])
def test_batch_that_completed_before_cleanup_deletes_output_and_error_files(
self, output_delete_fails: bool
) -> None:
client: Final = CleanupClient(
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=(
UnknownApiError(status_code=403, body="forbidden") if output_delete_fails else deleted_file(),
deleted_file(),
),
)
if output_delete_fails:
with pytest.raises(ExceptionGroup, match="output cleanup failed"):
cleanup_batch(client, "batch-1", key="test-key", provider="openai", delete_output_files=True)
else:
cleanup_batch(client, "batch-1", key="test-key", provider="openai", delete_output_files=True)
client.calls.assert_done()
@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(("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")
else:
with pytest.raises(AssertionError, match="Cancel batch batch-1 left status in_progress"):
cleanup_batch(client, "batch-1", key="test-key")
client.calls.assert_done()
class TestAzureFileExpiry:
def test_azure_form_serializes_native_expiry_for_the_proxy(self) -> None:
form: Final = batch_upload_form("azure", target_model_names="azure-test")
assert form.model_dump(by_alias=True, exclude_none=True) == {
"purpose": "batch",
"target_model_names": "azure-test",
"expires_after[anchor]": "created_at",
"expires_after[seconds]": AZURE_FILE_EXPIRY_SECONDS,
}
@pytest.mark.parametrize("provider", ["openai", "vertex_ai", "bedrock"])
def test_other_providers_keep_their_existing_upload_fields(self, provider: str) -> None:
assert batch_upload_form(provider).model_dump(by_alias=True, exclude_none=True) == {"purpose": "batch"}

View file

@ -21,14 +21,16 @@ import os
import re
import time
from datetime import datetime, timedelta, timezone
from typing import Callable
import pytest
from pydantic import BaseModel
from e2e_config import PROXY_BASE_URL, unique_marker
from e2e_config import MASTER_KEY, PROXY_BASE_URL, unique_marker
from batch_cleanup import cleanup_batch, cleanup_file
from batch_client import (
AZURE_FILE_EXPIRY_SECONDS,
batch_upload_form,
UPLOAD_FILENAME,
BatchClient,
BatchCreateBody,
@ -155,19 +157,19 @@ def upload_for_scenario(
if cap.scenario == "encoded":
return client.upload_file(
content=content,
form=FileUploadForm(purpose="batch"),
form=batch_upload_form(cap.provider),
model=cap.model,
key=key,
)
if cap.scenario == "unified":
return client.upload_file(
content=content,
form=FileUploadForm(purpose="batch", target_model_names=cap.model),
form=batch_upload_form(cap.provider, target_model_names=cap.model),
key=key,
)
return client.upload_file(
content=content,
form=FileUploadForm(purpose="batch"),
form=batch_upload_form(cap.provider),
key=key,
provider=cap.provider,
)
@ -188,20 +190,11 @@ def create_for_scenario(
def op_provider(cap: Capability) -> str | None:
"""provider_fallback ids are raw, so retrieve/cancel/list/delete need the provider
"""provider_fallback batch ids are raw, so retrieve/cancel/list need the provider
hint; the other scenarios encode it into the id and route automatically."""
return cap.provider if cap.scenario == "provider_fallback" else None
def quietly(action: Callable[[], object]) -> Callable[[], None]:
"""Adapt a value-returning call into a best-effort cleanup the teardown can run."""
def run() -> None:
action()
return run
def assert_file_object(file: FileObject, *, provider: str) -> None:
assert file.object == "file", f"file.object={file.object!r}"
assert file.purpose == "batch", f"file.purpose={file.purpose!r}"
@ -209,6 +202,10 @@ def assert_file_object(file: FileObject, *, provider: str) -> None:
if provider != "bedrock":
assert file.bytes > 0, f"file.bytes={file.bytes!r}"
assert file.status, "file.status missing"
if provider == "azure":
assert file.expires_at is not None, "Azure batch input has no automatic expiry"
assert file.created_at is not None
assert file.expires_at - file.created_at == AZURE_FILE_EXPIRY_SECONDS
assert (
file.created_at is not None and file.created_at > 0
), "file.created_at missing"
@ -249,7 +246,7 @@ def test_batch_lifecycle(
file = unwrap(upload_for_scenario(client, cap, render_jsonl(cap.jsonl_model), key))
resources.defer(
quietly(lambda: client.delete_file(file.id, key=key, provider=provider))
lambda: cleanup_file(client, file.id, key=key, provider=cap.file_provider)
)
assert_file_object(file, provider=cap.provider)
assert matches_id_shape(
@ -260,7 +257,9 @@ def test_batch_lifecycle(
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(
quietly(lambda: client.cancel_batch(batch.id, key=key, provider=provider))
lambda: cleanup_batch(
client, batch.id, key=key, provider=provider, delete_output_files=cap.provider in {"openai", "azure"}
)
)
assert batch.id, f"create returned no batch id (body={created.body[:200]})"
@ -339,7 +338,7 @@ def test_batch_key_model_access_denied(
denied_upload = client.upload_file(
content=render_jsonl(AZURE_BATCH_MODEL),
form=FileUploadForm(purpose="batch"),
form=batch_upload_form("azure"),
model=AZURE_BATCH_MODEL,
key=key,
)
@ -356,7 +355,7 @@ def test_batch_key_model_access_denied(
)
).id
resources.defer(
quietly(lambda: client.delete_file(raw_file, key=key, provider="openai"))
lambda: cleanup_file(client, raw_file, key=key, provider="openai")
)
denied_create = client.create_batch(
@ -383,6 +382,7 @@ def test_file_upload_and_delete_outputs(
key=key,
)
)
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="openai")
deleted = unwrap(client.delete_file(file.id, key=key))
@ -458,12 +458,12 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
_ = client.proxy.poll_logs_for_key(key, min_rows=1)
@ -517,7 +517,7 @@ class TestBatchFileContent:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert file.id
downloaded = client.proxy.transport.download(
@ -559,11 +559,11 @@ class TestBatchFileContent:
file = unwrap(
client.upload_file(
content=payload,
form=FileUploadForm(purpose="batch", target_model_names=provider.model),
form=batch_upload_form(provider.name, target_model_names=provider.model),
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider=provider.name)
assert is_managed_id(file.id), (
f"{provider.name}: unified upload must return a managed file id, got {file.id!r}"
@ -626,7 +626,7 @@ class TestOpenAIFiles:
)
)
resources.defer(
quietly(lambda: client.delete_file(file.id, key=key, provider="openai"))
lambda: cleanup_file(client, file.id, key=key, provider="openai")
)
listed = unwrap(client.list_files(key=key))
@ -690,7 +690,7 @@ class TestOpenAIFiles:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
fetched = unwrap(client.retrieve_file(file.id, key=key))
assert fetched.id == file.id, "retrieve must echo the uploaded file id"
@ -760,7 +760,7 @@ class TestBatchRateLimitErrorMapping:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
@ -803,7 +803,7 @@ class TestBatchEnqueuedTokenLimit:
"""
def _upload_batch_file(
self, client: BatchClient, resources: ResourceManager, key: str
self, client: BatchClient, resources: ResourceManager, key: str, *, cleanup_key: str | None = None
) -> FileObject:
file = unwrap(
client.upload_file(
@ -813,7 +813,7 @@ class TestBatchEnqueuedTokenLimit:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=cleanup_key or key))
return file
def _generate_enqueued_key(
@ -850,7 +850,7 @@ class TestBatchEnqueuedTokenLimit:
marker="rpm",
rpm_limit=BATCH_RL_RPM_LIMIT,
)
file = self._upload_batch_file(client, resources, key)
file = self._upload_batch_file(client, resources, key, cleanup_key=MASTER_KEY)
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
@ -861,7 +861,7 @@ class TestBatchEnqueuedTokenLimit:
)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=MASTER_KEY, delete_output_files=True))
@pytest.mark.covers(
"quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted",
@ -904,7 +904,7 @@ class TestBatchEnqueuedTokenLimit:
first = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(first)
first_batch = BatchObject.model_validate_json(first.body)
resources.defer(quietly(lambda: client.cancel_batch(first_batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, first_batch.id, key=key))
blocked = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert blocked.status_code == 429, (
@ -928,7 +928,7 @@ class TestBatchEnqueuedTokenLimit:
)
require_successful_call(retried)
retry_batch = BatchObject.model_validate_json(retried.body)
resources.defer(quietly(lambda: client.cancel_batch(retry_batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, retry_batch.id, key=key))
ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
@ -984,13 +984,13 @@ class TestBedrockBatchAssumeRole:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="bedrock")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert batch.id, f"assume-role create returned no batch id: {created.body[:200]}"
assert is_managed_id(batch.id), (
@ -1044,7 +1044,7 @@ class TestGeminiFiles:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="gemini")
assert file.id, "gemini file upload returned no id"
@ -1099,13 +1099,13 @@ class TestHostedVllmBatch:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="hosted_vllm")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert batch.id, f"hosted_vllm create returned no batch id: {created.body[:200]}"
assert batch.status in CREATED_BATCH_STATUSES, (
@ -1192,7 +1192,7 @@ class TestBatchFailurePaths:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
@ -1243,12 +1243,12 @@ class TestBatchFailurePaths:
file = unwrap(
client.upload_file(
content=render_jsonl(AZURE_BATCH_RAW_MODEL),
form=FileUploadForm(purpose="batch"),
form=batch_upload_form("azure"),
model=AZURE_BATCH_MODEL,
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert decoded_model_from_id(file.id) == AZURE_BATCH_MODEL, (
f"upload did not encode the azure deployment into the file id: {file.id!r}"
)
@ -1258,7 +1258,7 @@ class TestBatchFailurePaths:
)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert decoded_model_from_id(batch.id) == AZURE_BATCH_MODEL, (
"create with a foreign encoded file id must route by the file's embedded model, "
@ -1307,7 +1307,7 @@ class TestBatchSecondHop:
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert is_managed_id(file.id), (
f"second-hop unified upload must return a managed file id, got {file.id!r}"
)
@ -1315,7 +1315,7 @@ class TestBatchSecondHop:
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert is_managed_id(batch.id), (
f"second-hop create must return a managed batch id, got {batch.id!r}"

View file

@ -21,6 +21,7 @@ from typing import Iterator
import pytest
from batch_client import BatchClient, FileObject
from batch_cleanup import cleanup_file
from capabilities import batch_model_name, is_managed_id, openai_batch_params
from e2e_config import unique_marker
from e2e_http import FileUploadForm, Result, UnknownApiError, unwrap
@ -108,7 +109,7 @@ def test_cross_user_managed_id_denied_owner_allowed(
key=owner_key,
)
)
resources.defer(lambda: client.delete_file(uploaded.id, key=owner_key))
resources.defer(lambda: cleanup_file(client, uploaded.id, key=owner_key))
assert is_managed_id(uploaded.id), f"expected a managed unified file id, got {uploaded.id}"
denied = client.retrieve_file(uploaded.id, key=other_key)

View file

@ -7,7 +7,7 @@ from __future__ import annotations
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import Literal
from typing import Final, Literal
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, settle_propagation, unique_marker
from e2e_http import NoBody, Result, StreamingResponse, Success, unwrap
@ -405,6 +405,29 @@ def build_client(proxy: ProxyClient) -> GuardrailsClient:
return GuardrailsClient(proxy=proxy)
def poll_until_guardrail_applied(
call: Callable[[], StreamingResponse],
guardrail_name: str,
*,
timeout: float = POLL_TIMEOUT,
interval: float = POLL_INTERVAL,
now: Callable[[], float] = time.monotonic,
sleep: Callable[[float], None] = time.sleep,
) -> StreamingResponse:
deadline: Final = now() + timeout
if not (result := call()).ok:
return result
while (
guardrail_name
not in (name.strip() for name in result.headers.get("x-litellm-applied-guardrails", "").split(","))
and (remaining := deadline - now()) > 0
):
sleep(min(interval, remaining))
if now() >= deadline or not (result := call()).ok:
break
return result
def poll_until_blocked[R: BaseModel](call: Callable[[], Result[R]]) -> Result[R]:
"""Retry a call that a guardrail should reject until it is, returning the last result.

View file

@ -0,0 +1,66 @@
from dataclasses import dataclass
from itertools import chain, repeat
from typing import Final
import pytest
from e2e_http import StreamingResponse
from guardrails_client import poll_until_guardrail_applied
@dataclass
class Clock:
elapsed: float = 0.0
def now(self) -> float:
return self.elapsed
def sleep(self, seconds: float) -> None:
self.elapsed += seconds
def _response(applied: str, status: int = 200) -> StreamingResponse:
return StreamingResponse(status_code=status, body="{}", headers={"x-litellm-applied-guardrails": applied})
def test_waits_for_requested_guardrail_after_an_unrelated_global_guardrail() -> None:
clock: Final = Clock()
expected: Final = _response("global-filter, tool-permission")
responses: Final = iter((_response("global-filter"), expected))
result: Final = poll_until_guardrail_applied(
lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
)
assert result is expected
assert clock.elapsed == 2
@pytest.mark.parametrize("applied", ("", "global-filter", "tool-permission-sibling"))
def test_missing_exact_guardrail_returns_failure_evidence_at_deadline(applied: str) -> None:
clock: Final = Clock()
missing: Final = _response(applied)
responses: Final = iter((missing, missing, missing))
result: Final = poll_until_guardrail_applied(
lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
)
assert result is missing
assert clock.elapsed == 5
with pytest.raises(StopIteration):
next(responses)
@pytest.mark.parametrize("status", (400, 401, 429, 500))
def test_http_failure_is_not_hidden_by_a_later_success(status: int) -> None:
clock: Final = Clock()
failed: Final = _response("", status)
responses: Final = iter(chain((failed,), repeat(_response("tool-permission"))))
result: Final = poll_until_guardrail_applied(
lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
)
assert result is failed
assert clock.elapsed == 0

View file

@ -30,6 +30,7 @@ from guardrails_client import (
ToolPermissionParamsBody,
ToolPermissionRuleBody,
poll_until_blocked,
poll_until_guardrail_applied,
)
from lifecycle import ResourceManager
from models import ChatResponse, ChatTool, ChatToolFunction
@ -84,8 +85,8 @@ def _register_tool_permission(client: GuardrailsClient, resources: ResourceManag
resources.defer(lambda: client.delete_guardrail(guardrail_id))
def _applied_guardrails(outcome: StreamingResponse) -> str:
return outcome.headers.get("x-litellm-applied-guardrails", "")
def _applied_guardrails(outcome: StreamingResponse) -> tuple[str, ...]:
return tuple(name.strip() for name in outcome.headers.get("x-litellm-applied-guardrails", "").split(","))
def _tool_call_names(response: ChatResponse) -> tuple[str, ...]:
@ -144,14 +145,17 @@ class TestToolPermissionPreCall:
name = f"e2e-toolperm-allow-{unique_marker()}"
_register_tool_permission(client, resources, name=name)
outcome = client.chat_raw(
scoped_key,
MODEL,
TOOL_PROMPT,
guardrails=[name],
max_tokens=128,
tools=[ALLOWED_TOOL],
tool_choice="required",
outcome = poll_until_guardrail_applied(
lambda: client.chat_raw(
scoped_key,
MODEL,
TOOL_PROMPT,
guardrails=[name],
max_tokens=128,
tools=[ALLOWED_TOOL],
tool_choice="required",
),
name,
)
assert outcome.ok, f"the permitted tool must be served, got {outcome.status_code}: {outcome.body[:400]}"

View file

@ -8,8 +8,9 @@ ResourceManager; the test registers a cleanup for every resource it creates, and
the fixture's teardown releases them all even when the test body raises.
"""
from builtins import ExceptionGroup
from dataclasses import dataclass, field
from typing import Callable, List, Protocol, runtime_checkable
from typing import Callable, Final, List, Protocol, runtime_checkable
from proxy_client import ProxyClient
from models import KeyGenerateBody
@ -52,6 +53,7 @@ class ResourceManager:
"""
client: ResourceClient
strict_cleanup: bool = False
_cleanups: List[Callable[[], object]] = field(
default_factory=list
) # mutable-ok: append-only teardown registry
@ -82,8 +84,17 @@ class ResourceManager:
return customer_id
def teardown(self) -> None:
for cleanup in reversed(self._cleanups):
try:
cleanup()
except Exception:
pass # best-effort: a failed cleanup must not block the rest
failures: Final = tuple(
failure for cleanup in reversed(self._cleanups)
if (failure := _run_cleanup(cleanup)) is not None
)
if failures and self.strict_cleanup:
raise ExceptionGroup("Resource cleanup failed", failures)
def _run_cleanup(cleanup: Callable[[], object]) -> Exception | None:
try:
cleanup()
except Exception as exc:
return exc
return None

View file

@ -1095,6 +1095,110 @@ async def test_afile_content_passes_trusted_model_credentials_to_router():
assert trusted_credentials["s3_bucket_name"] == "my-bucket"
def _managed_deletion_file_id(provider_file_id):
from litellm.types.utils import SpecialEnums
value = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/json", "test-file", "batch-model", provider_file_id, "model-123"
)
return base64.urlsafe_b64encode(value.encode()).decode().rstrip("=")
def _managed_files_with_deletion_row(unified_file_id, provider_file_id, file_object):
from litellm.caching import DualCache
from litellm.models.managed_files import LiteLLM_ManagedFileTable
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
row = LiteLLM_ManagedFileTable(
unified_file_id=unified_file_id,
model_mappings={"model-123": provider_file_id},
flat_model_file_ids=[provider_file_id],
file_object=file_object,
)
table = MagicMock(
find_first=AsyncMock(return_value=row),
delete=AsyncMock(),
)
return _PROXY_LiteLLMManagedFiles(
internal_usage_cache=DualCache(),
prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=table)),
), table
@pytest.mark.asyncio
async def test_afile_delete_bedrock_uses_deployment_bucket_and_signed_s3_delete(monkeypatch):
import httpx
import respx
from litellm import Router
monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
monkeypatch.delenv("AWS_S3_OUTPUT_BUCKET_NAME", raising=False)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
router = Router(
model_list=[
{
"model_name": "bedrock-batch",
"litellm_params": {
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"aws_access_key_id": "AKIAEXAMPLE",
"aws_secret_access_key": "secret",
"aws_region_name": "us-west-2",
"s3_bucket_name": "my-bucket",
},
"model_info": {"id": "model-123"},
}
],
num_retries=0,
)
s3_uri = "s3://my-bucket/litellm-bedrock-files/input.jsonl"
unified_file_id = _managed_deletion_file_id(s3_uri)
managed_files, table = _managed_files_with_deletion_row(unified_file_id, s3_uri, None)
with respx.mock:
route = respx.delete(
"https://s3.us-west-2.amazonaws.com/my-bucket/litellm-bedrock-files/input.jsonl"
).mock(return_value=httpx.Response(204))
response = await managed_files.afile_delete(
file_id=unified_file_id,
litellm_parent_otel_span=None,
llm_router=router,
_litellm_internal_model_credentials={"s3_bucket_name": "request-bucket"},
)
assert len(route.calls) == 1
assert route.calls[0].request.headers["Authorization"].startswith("AWS4-HMAC-SHA256")
assert response.id == unified_file_id
assert response.deleted is True
table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
@pytest.mark.asyncio
async def test_afile_delete_returns_managed_id_for_stored_provider_output():
from openai.types import FileDeleted
provider_file_id = "file-error-output"
unified_file_id = _managed_deletion_file_id(provider_file_id)
stored_file = _make_file_object(provider_file_id)
managed_files, table = _managed_files_with_deletion_row(unified_file_id, provider_file_id, stored_file)
router = MagicMock(
get_deployment_credentials_with_provider=MagicMock(return_value=None),
afile_delete=AsyncMock(return_value=FileDeleted(id=provider_file_id, object="file", deleted=True)),
)
response = await managed_files.afile_delete(
file_id=unified_file_id,
litellm_parent_otel_span=None,
llm_router=router,
_litellm_internal_model_credentials={"s3_bucket_name": "request-bucket"},
)
assert response.id == unified_file_id
assert response.object == "file"
assert response.filename == stored_file.filename
assert stored_file.id == provider_file_id
router.afile_delete.assert_awaited_once_with(model="model-123", file_id=provider_file_id)
table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
@pytest.mark.asyncio
async def test_afile_content_bedrock_unified_id_end_to_end(monkeypatch):
"""

View file

@ -5,6 +5,8 @@ Test bedrock files transformation functionality
import json
import os
from collections.abc import Mapping
from contextlib import AsyncExitStack, closing
from typing import Final
from unittest.mock import MagicMock
from urllib.parse import unquote, urlparse
@ -1855,6 +1857,104 @@ class TestBedrockBatchNonChatEndpointRecords:
]
class TestBedrockFileDeletion:
S3_URI: Final = "s3://my-bucket/litellm-bedrock-files-model-abc.jsonl"
URL: Final = "https://s3.us-west-2.amazonaws.com/my-bucket/litellm-bedrock-files-model-abc.jsonl"
def test_interleaved_deletions_keep_their_own_file_ids(self, monkeypatch: pytest.MonkeyPatch) -> None:
import httpx
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
config: Final = BedrockFilesConfig()
params: Final = {
"aws_access_key_id": "AKIAEXAMPLE",
"aws_secret_access_key": "test-secret",
"aws_region_name": "us-west-2",
}
file_ids: Final = (self.S3_URI, "s3://my-bucket/litellm-bedrock-files-model-second.jsonl")
for file_id in file_ids:
config.transform_delete_file_request(file_id=file_id, optional_params={}, litellm_params=params)
deleted: Final = tuple(
config.transform_delete_file_response(
raw_response=httpx.Response(204),
logging_obj=MagicMock(model_call_details={"additional_args": {"file_id": file_id}}),
litellm_params=params,
).id
for file_id in file_ids
)
assert deleted == file_ids
def test_delete_file_sends_signed_delete_and_returns_matching_id(self, monkeypatch: pytest.MonkeyPatch) -> None:
import httpx
import respx
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
with respx.mock, closing(HTTPHandler()) as client:
route: Final = respx.delete(self.URL).mock(return_value=httpx.Response(204))
deleted: Final = litellm.file_delete(
file_id=self.S3_URI, custom_llm_provider="bedrock", client=client,
aws_access_key_id="AKIAEXAMPLE", aws_secret_access_key="test-secret", aws_region_name="us-west-2",
)
assert route.call_count == 1
request: Final = route.calls[0].request
assert request.content == b""
signed: Final = AWSRequest(method="DELETE", url=self.URL, headers={
"X-Amz-Date": request.headers["X-Amz-Date"],
"X-Amz-Content-SHA256": request.headers["X-Amz-Content-SHA256"],
})
signed.context["timestamp"] = request.headers["X-Amz-Date"]
auth: Final = S3SigV4Auth(Credentials("AKIAEXAMPLE", "test-secret"), "s3", "us-west-2")
signature: Final = auth.signature(auth.string_to_sign(signed, auth.canonical_request(signed)), signed)
assert request.headers["Authorization"].endswith(f"Signature={signature}")
assert deleted.id == self.S3_URI and deleted.deleted is True
@pytest.mark.asyncio
async def test_adelete_file_propagates_s3_errors(self, monkeypatch: pytest.MonkeyPatch) -> None:
import httpx
import respx
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
async with AsyncExitStack() as stack:
client: Final = AsyncHTTPHandler()
stack.push_async_callback(client.close)
with respx.mock:
route: Final = respx.delete(self.URL).mock(
return_value=httpx.Response(403, content=b"<Error><Code>AccessDenied</Code></Error>")
)
from litellm.llms.bedrock.common_utils import BedrockError
with pytest.raises(BedrockError, match="AccessDenied"):
await litellm.afile_delete(
file_id=self.S3_URI, custom_llm_provider="bedrock", client=client,
aws_access_key_id="AKIAEXAMPLE", aws_secret_access_key="test-secret", aws_region_name="us-west-2",
)
assert route.call_count == 1
@pytest.mark.parametrize("file_id, message", [
("s3://other-bucket/litellm-bedrock-files-model-abc.jsonl", "configured storage bucket"),
("s3://my-bucket/private/data.jsonl", "LiteLLM-managed"),
])
def test_delete_rejects_untrusted_objects_before_signing(
self, file_id: str, message: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
with pytest.raises(ValueError, match=message):
BedrockFilesConfig().transform_delete_file_request(file_id=file_id, optional_params={}, litellm_params={})
class TestBedrockFileContentTransformation:
"""SigV4-signed S3 GetObject retrieval of Bedrock batch output files."""
@ -1873,7 +1973,7 @@ class TestBedrockFileContentTransformation:
import hashlib
from litellm.llms.bedrock.files.transformation import (
S3_SIGNED_GET_HEADERS_PARAM,
S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@ -1889,7 +1989,7 @@ class TestBedrockFileContentTransformation:
assert url == self.EXPECTED_URL
assert params == {}
signed_headers = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]
signed_headers = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]
content_hashes = {
value
for name, value in signed_headers.items()
@ -2139,7 +2239,7 @@ class TestBedrockFileContentTransformation:
def test_s3_region_name_wins_for_content_signing(self, monkeypatch):
"""s3_region_name must override aws_region_name for both the URL and the signature."""
from litellm.llms.bedrock.files.transformation import (
S3_SIGNED_GET_HEADERS_PARAM,
S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@ -2154,17 +2254,17 @@ class TestBedrockFileContentTransformation:
)
assert url.startswith("https://s3.eu-west-1.amazonaws.com/")
authorization = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]["Authorization"]
authorization = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Authorization"]
assert "/eu-west-1/s3/aws4_request" in authorization
def test_validate_environment_merges_and_pops_signed_get_headers(self):
from litellm.llms.bedrock.files.transformation import (
S3_SIGNED_GET_HEADERS_PARAM,
S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
litellm_params = {
S3_SIGNED_GET_HEADERS_PARAM: {"Authorization": "AWS4-HMAC-SHA256 test"}
S3_SIGNED_REQUEST_HEADERS_PARAM: {"Authorization": "AWS4-HMAC-SHA256 test"}
}
headers = BedrockFilesConfig().validate_environment(
@ -2179,7 +2279,7 @@ class TestBedrockFileContentTransformation:
"x-custom": "kept",
"Authorization": "AWS4-HMAC-SHA256 test",
}
assert S3_SIGNED_GET_HEADERS_PARAM not in litellm_params
assert S3_SIGNED_REQUEST_HEADERS_PARAM not in litellm_params
def test_transform_file_content_response_wraps_binary_content(self):
import httpx
@ -2379,7 +2479,7 @@ class TestBedrockFilesS3SignatureEncoding:
self, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.llms.bedrock.files.transformation import (
S3_SIGNED_GET_HEADERS_PARAM,
S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@ -2402,7 +2502,7 @@ class TestBedrockFilesS3SignatureEncoding:
method="GET",
url=url,
body=None,
headers=litellm_params[S3_SIGNED_GET_HEADERS_PARAM],
headers=litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM],
)
@ -2457,7 +2557,7 @@ def test_sign_s3_request_assumes_role_with_external_id(monkeypatch):
assert "ASIAFILESPUTROLE" in authorization
def test_sign_s3_get_request_assumes_role_with_external_id(monkeypatch):
def test_sign_s3_request_without_body_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied when signing the S3 download request."""
import datetime
from unittest.mock import patch
@ -2504,7 +2604,7 @@ def test_sign_s3_get_request_assumes_role_with_external_id(monkeypatch):
assert request_params.aws_external_id == "external-id-files-get"
with patch.object(boto3, "client", return_value=FakeSTSClient()):
signed_headers = BedrockFilesConfig()._sign_s3_get_request(
signed_headers = BedrockFilesConfig()._sign_s3_request_without_body(
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
aws_region_name="us-east-1",
request_params=request_params,