mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge pull request #40161 from BerriAI/litellm_batch_e2e_cleanup
fix(batches): clean up E2E resources across providers
This commit is contained in:
commit
802e526cf9
17 changed files with 949 additions and 108 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
140
tests/e2e/batches/batch_cleanup.py
Normal file
140
tests/e2e/batches/batch_cleanup.py
Normal 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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
313
tests/e2e/batches/test_batch_cleanup.py
Normal file
313
tests/e2e/batches/test_batch_cleanup.py
Normal 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"}
|
||||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
66
tests/e2e/guardrails/test_guardrails_client.py
Normal file
66
tests/e2e/guardrails/test_guardrails_client.py
Normal 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
|
||||
|
|
@ -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]}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue