feat(proxy): cap batch file records, daily batch uploads, and per-file downloads (#43632)

* feat(proxy): cap batch file records, daily batch uploads, and per-file downloads

Adds three opt-in limits for batch jobs, each settable in general_settings as a
per-key default and overridable in key or team metadata by a proxy admin:

max_batch_file_records rejects a purpose=batch upload with more request lines
than allowed with a 413 before it reaches the provider.

max_batch_file_uploads_per_day counts accepted batch uploads per key and per
team in a UTC day and returns 429 with Retry-After once the count is used.

max_file_downloads_per_minute counts GET /v1/files/{id}/content per key and per
team for each file in a one-minute window and returns 429 with Retry-After.

The existing admin-only guard for batch_enqueued_token_limit now covers all four
metadata keys.

* fix(proxy): keep file usage counters in their own store and gate batch limits on user creation

* fix(proxy): count keyless JWT callers per user, require positive file caps, and take the upload slot after request validation

* refactor(proxy): end file usage cap describers with an explicit return after the match

* fix(proxy): keep file usage counters when more than 200 are live without Redis

The file usage counter store used a default in-memory cache, which holds 200
entries and evicts the one that expires soonest. Without Redis, a caller got a
fresh per-file download allowance after touching about 200 other file ids in
the same minute, and a key got a fresh daily upload allowance once about 200
other keys had uploaded that day. The store now tracks up to 20,000 live
counters per worker, the same bound the login throttle uses

* test(proxy): move the file usage cap tests into the directory the proxy shard runs

Main's shard coverage check found tests/unit/proxy/openai_files_endpoints
claimed by no shard, so its tests would not run in CI. The file moves next to
the other files endpoint tests in tests/unit/proxy/openai_files_endpoint, which
the proxy-endpoints shard already runs

* fix(proxy): declare the file usage counters as rate limit calls

Main's redis producer gate requires every module that writes a shared cache to name its key family, and the file usage counters wrote theirs without one.

* test(files): audit batch file usage caps across processes, Redis outages, and config reloads

* test(files): guard the chaos cells against minute boundaries and open Redis breakers

Two chaos cells each failed once in the audit run. The restart check ran
three sequential downloads with no guard against straddling a UTC minute,
and the exact-cap probe after a Redis outage ran while both workers' Redis
circuit breakers were still open (60 s default recovery), so it counted in
per-process memory and the two workers split the cap

Every burst now carries a window guard, the chaos fixture lowers the
breaker recovery to 2 s, and the post-outage check drives a fresh key to its
cap through a one-worker sibling proxy and then expects the two-worker
candidate to refuse the whole burst, which only the shared Redis count can
produce, polled until the breakers close

* test(files): release the held uploads when the killed-worker cell fails early

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 03:06:36 +00:00 • committed by GitHub
parent 12dcce8db2
commit 402fa63366
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 3134 additions and 51 deletions

View file

@ -2209,6 +2209,16 @@ BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60
# admins may write it: when present it replaces the standard RPM/TPM checks for
# batch submissions.
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit"
MAX_BATCH_FILE_RECORDS_KEY: Final = "max_batch_file_records"
MAX_BATCH_FILE_UPLOADS_PER_DAY_KEY: Final = "max_batch_file_uploads_per_day"
MAX_FILE_DOWNLOADS_PER_MINUTE_KEY: Final = "max_file_downloads_per_minute"
ADMIN_ONLY_BATCH_LIMIT_METADATA_KEYS: Final = (
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
MAX_BATCH_FILE_RECORDS_KEY,
MAX_BATCH_FILE_UPLOADS_PER_DAY_KEY,
MAX_FILE_DOWNLOADS_PER_MINUTE_KEY,
)
FILE_USAGE_MAX_TRACKED_COUNTERS: Final = 20_000
# Shared read-only empty mapping, for defaulting optional Mapping parameters without
# constructing a fresh mutable dict at each call site.

View file

@ -2891,6 +2891,21 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider",
)
max_batch_file_records: int | None = Field(
None,
gt=0,
description="max records (non-blank lines) per batch input file for /v1/files uploads with purpose=batch, applied per key. A key's metadata can override it and a team's metadata adds a team cap on top, both set by a proxy admin; the lower of the key's value and the team's value wins. Unset means no limit",
)
max_batch_file_uploads_per_day: int | None = Field(
None,
gt=0,
description="max /v1/files uploads with purpose=batch per key (per user for JWT callers) per UTC day. A key's metadata can override it and a team's metadata adds a shared team cap, both set by a proxy admin. Unset means no limit",
)
max_file_downloads_per_minute: int | None = Field(
None,
gt=0,
description="max GET /v1/files/{file_id}/content calls per key (per user for JWT callers) per file per minute. A key's metadata can override it and a team's metadata adds a shared team cap, both set by a proxy admin. Unset means no limit",
)
max_file_size_mb: int | None = Field(
None,
description="max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider",

View file

@ -14,7 +14,7 @@ import litellm
from litellm import Router, constants, provider_list
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
ADMIN_ONLY_BATCH_LIMIT_METADATA_KEYS,
EMPTY_MAPPING,
INVALID_VIRTUAL_KEY_ERROR_MARKER,
MINIMUM_CUSTOM_KEY_LENGTH,
@ -1324,8 +1324,8 @@ def enforce_output_token_estimates_are_admin_only(
)
class BatchEnqueuedTokenLimitRequest(Protocol):
"""The shape of any management request that can carry a batch enqueued-token limit."""
class BatchLimitRequest(Protocol):
"""The shape of any management request that can carry an admin-only batch limit in its metadata."""
@property
def metadata(self) -> Mapping[str, object] | None: ...
@ -1334,18 +1334,18 @@ class BatchEnqueuedTokenLimitRequest(Protocol):
def model_fields_set(self) -> Collection[str]: ...
def enforce_batch_enqueued_token_limit_is_admin_only(
data: BatchEnqueuedTokenLimitRequest,
def enforce_batch_limits_are_admin_only(
data: BatchLimitRequest,
existing_metadata: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
entity: Literal["key", "team"],
) -> None:
"""Only a proxy admin may change a key or team's batch enqueued-token limit.
"""Only a proxy admin may change a key or team's batch limits.
When set, ``batch_enqueued_token_limit`` replaces the standard RPM/TPM checks
for batch submissions, so a holder-writable copy would let a caller lift their
own batch quota. Gated on the resulting value rather than on presence, so a
form resending the stored value stays a no-op.
Every key in ``ADMIN_ONLY_BATCH_LIMIT_METADATA_KEYS`` caps what the holder
may do with batches, so a holder-writable copy would let a caller lift their
own quota. Gated on the resulting value rather than on presence, so a form
resending the stored value stays a no-op.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
@ -1353,13 +1353,17 @@ def enforce_batch_enqueued_token_limit_is_admin_only(
requested: Final[Mapping[str, object]] = (
(data.metadata or EMPTY_MAPPING) if "metadata" in data.model_fields_set else stored
)
if requested.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) == stored.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY):
changed: Final = next(
(key for key in ADMIN_ONLY_BATCH_LIMIT_METADATA_KEYS if requested.get(key) != stored.get(key)),
None,
)
if changed is None:
return
raise HTTPException(
status_code=403,
detail={
"error": f"Only proxy admins can set {BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY} on a {entity}. "
"It replaces the standard rate limit checks for batch submissions."
"error": f"Only proxy admins can set {changed} on a {entity}. "
"It limits what the holder can do with batches, so the holder cannot raise it."
},
)

View file

@ -37,6 +37,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_object,
get_user_object,
)
from litellm.proxy.auth.auth_utils import enforce_batch_limits_are_admin_only
from litellm.proxy.auth.password_policy import (
validate_password_not_breached,
validate_password_policy,
@ -587,6 +588,9 @@ async def new_user(
detail=f"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). Attempted to create user with role: {data.user_role}. Your role: {user_api_key_dict.user_role}",
)
if data.auto_create_key and isinstance(user_api_key_dict, UserAPIKeyAuth):
enforce_batch_limits_are_admin_only(data, None, user_api_key_dict, "key")
_check_permissions_caller_permission(
data=data,
user_api_key_dict=user_api_key_dict,

View file

@ -64,7 +64,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
enforce_batch_enqueued_token_limit_is_admin_only,
enforce_batch_limits_are_admin_only,
enforce_output_token_estimates_are_admin_only,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -1218,7 +1218,7 @@ async def _common_key_generation_helper(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=data,
existing_metadata=None,
user_api_key_dict=user_api_key_dict,
@ -2757,7 +2757,7 @@ async def _process_single_key_update(
existing_metadata=existing_key_row.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=update_key_request,
existing_metadata=existing_key_row.metadata,
user_api_key_dict=user_api_key_dict,
@ -3205,7 +3205,7 @@ async def _validate_update_key_data(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=data,
existing_metadata=_existing_metadata if isinstance(_existing_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,
@ -5598,7 +5598,7 @@ async def _execute_virtual_key_regeneration(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=data,
existing_metadata=_existing_key_metadata if isinstance(_existing_key_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,

View file

@ -109,7 +109,7 @@ from litellm.proxy.auth.auth_checks import (
invalidate_team_member_spend_state,
)
from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
enforce_batch_limits_are_admin_only,
enforce_output_token_estimates_are_admin_only,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -1488,7 +1488,7 @@ async def new_team(
user_api_key_dict=user_api_key_dict,
entity="team",
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=data,
existing_metadata=None,
user_api_key_dict=user_api_key_dict,
@ -2274,7 +2274,7 @@ async def update_team(
user_api_key_dict=user_api_key_dict,
entity="team",
)
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=data,
existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,

View file

@ -29,6 +29,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
from litellm.proxy.auth.auth_utils import enforce_batch_limits_are_admin_only
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
@ -213,6 +214,8 @@ def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str
try:
validate_budget_duration(item.budget_duration)
_check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict)
if item.auto_create_key:
enforce_batch_limits_are_admin_only(item, None, user_api_key_dict, "key")
except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only
return _error_message(exc)
return None

View file

@ -7,6 +7,7 @@ from typing import BinaryIO, Final, NoReturn
from typing_extensions import assert_never
from litellm.proxy._types import ProxyException
from litellm.proxy.openai_files_endpoints.file_usage_caps import FileUsageLimit, describe_limit_source
_MB: Final = 1024 * 1024
@ -33,6 +34,11 @@ class BatchFileTooLarge:
limit_mb: int
@dataclass(frozen=True, slots=True)
class BatchFileTooManyRecords:
limit: FileUsageLimit
@dataclass(frozen=True, slots=True)
class BatchFileWrongExtension:
filename: str
@ -62,6 +68,7 @@ class BatchFileMissingLineKey:
BatchFileValidationFailure = (
BatchFileTooLarge
| BatchFileTooManyRecords
| BatchFileWrongExtension
| BatchFileEmpty
| BatchFileInvalidJsonLine
@ -99,7 +106,23 @@ def _check_line(line_number: int, raw_line: bytes, line_shape: BatchLineShape) -
return BatchFileMissingLineKey(line_number=line_number, key=missing, line_shape=line_shape)
def _scan_lines(file_source: bytes | BinaryIO, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
def _check_record(
record_number: int,
line_number: int,
raw_line: bytes,
line_shape: BatchLineShape,
max_records: FileUsageLimit | None,
) -> BatchFileValidationFailure | None:
if max_records is not None and record_number > max_records.value:
return BatchFileTooManyRecords(limit=max_records)
return _check_line(line_number, raw_line, line_shape)
def _scan_lines(
file_source: bytes | BinaryIO,
line_shape: BatchLineShape,
max_records: FileUsageLimit | None,
) -> BatchFileValidationFailure | None:
content_lines: Final = (
(line_number, raw_line)
for line_number, raw_line in enumerate(_iter_lines(file_source), start=1)
@ -108,15 +131,11 @@ def _scan_lines(file_source: bytes | BinaryIO, line_shape: BatchLineShape) -> Ba
first_line: Final = next(content_lines, None)
if first_line is None:
return BatchFileEmpty()
return next(
(
failure
for line_number, raw_line in chain((first_line,), content_lines)
for failure in (_check_line(line_number, raw_line, line_shape),)
if failure is not None
),
None,
failures: Final = (
_check_record(record_number, line_number, raw_line, line_shape, max_records)
for record_number, (line_number, raw_line) in enumerate(chain((first_line,), content_lines), start=1)
)
return next((failure for failure in failures if failure is not None), None)
def check_batch_file_upload(
@ -124,6 +143,7 @@ def check_batch_file_upload(
file_source: bytes | BinaryIO,
max_batch_file_size_mb: int | None,
line_shape: BatchLineShape = BATCH_LINE_SHAPE,
max_records: FileUsageLimit | None = None,
) -> BatchFileValidationFailure | None:
if filename is None or not filename.lower().endswith(".jsonl"):
return BatchFileWrongExtension(filename=filename or "")
@ -131,7 +151,7 @@ def check_batch_file_upload(
size_bytes: Final = _file_size_bytes(file_source)
if size_bytes > max_batch_file_size_mb * _MB:
return BatchFileTooLarge(size_bytes=size_bytes, limit_mb=max_batch_file_size_mb)
scan_failure: Final = _scan_lines(file_source, line_shape)
scan_failure: Final = _scan_lines(file_source, line_shape, max_records)
if not isinstance(file_source, bytes):
file_source.seek(0)
return scan_failure
@ -149,6 +169,17 @@ def raise_batch_file_validation_failure(failure: BatchFileValidationFailure) ->
param="file",
code=413,
)
case BatchFileTooManyRecords(limit=limit):
raise ProxyException(
message=(
f"Batch input file has more than {limit.value} records, which exceeds the "
f"{limit.setting} of {limit.value} set {describe_limit_source(limit.source)}. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param="file",
code=413,
)
case BatchFileWrongExtension(filename=filename):
raise ProxyException(
message=(

View file

@ -0,0 +1,250 @@
import math
import time
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Annotated, Final, Literal, NoReturn, TypeAlias
from pydantic import Field, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
EMPTY_MAPPING,
MAX_BATCH_FILE_RECORDS_KEY,
MAX_BATCH_FILE_UPLOADS_PER_DAY_KEY,
MAX_FILE_DOWNLOADS_PER_MINUTE_KEY,
)
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache
FileUsageSetting: TypeAlias = Literal[
"max_batch_file_records",
"max_batch_file_uploads_per_day",
"max_file_downloads_per_minute",
]
LimitSource: TypeAlias = Literal["key", "team", "general_settings"]
CounterScope: TypeAlias = Literal["key", "user", "team"]
_COUNTER_PREFIX: Final = "litellm:file_usage"
_DAY_SECONDS: Final = 24 * 60 * 60
_MINUTE_SECONDS: Final = 60
_LIMIT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(gt=0)])
_METADATA_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
@dataclass(frozen=True, slots=True)
class FileUsageLimit:
setting: FileUsageSetting
value: int
source: LimitSource
@dataclass(frozen=True, slots=True)
class ScopedFileUsageLimit:
scope: CounterScope
scope_id: str
limit: FileUsageLimit
@dataclass(frozen=True, slots=True)
class FileUsageLimitExceeded:
limit: ScopedFileUsageLimit
retry_after_seconds: int
def _read_limit(
settings: Mapping[str, object],
setting: FileUsageSetting,
source: LimitSource,
) -> FileUsageLimit | None:
raw: Final = settings.get(setting)
if raw is None:
return None
try:
return FileUsageLimit(setting=setting, value=_LIMIT_ADAPTER.validate_python(raw), source=source)
except ValidationError:
verbose_proxy_logger.warning(
"Ignoring invalid %s in %s; expected a positive integer",
setting,
source,
)
return None
def _key_limit(
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
setting: FileUsageSetting,
) -> FileUsageLimit | None:
key_metadata: Final = _METADATA_ADAPTER.validate_python(user_api_key_dict.metadata or EMPTY_MAPPING)
from_key: Final = _read_limit(key_metadata, setting, "key")
if from_key is not None:
return from_key
return _read_limit(general_settings, setting, "general_settings")
def _team_limit(user_api_key_dict: UserAPIKeyAuth, setting: FileUsageSetting) -> FileUsageLimit | None:
team_metadata: Final = _METADATA_ADAPTER.validate_python(user_api_key_dict.team_metadata or EMPTY_MAPPING)
return _read_limit(team_metadata, setting, "team")
def batch_file_record_limit(
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
) -> FileUsageLimit | None:
applicable: Final = tuple(
limit
for limit in (
_key_limit(user_api_key_dict, general_settings, MAX_BATCH_FILE_RECORDS_KEY),
_team_limit(user_api_key_dict, MAX_BATCH_FILE_RECORDS_KEY),
)
if limit is not None
)
return min(applicable, key=lambda limit: limit.value, default=None)
def _caller_counter(user_api_key_dict: UserAPIKeyAuth, limit: FileUsageLimit | None) -> ScopedFileUsageLimit | None:
if limit is None:
return None
if user_api_key_dict.api_key:
return ScopedFileUsageLimit(scope="key", scope_id=user_api_key_dict.api_key, limit=limit)
if user_api_key_dict.user_id:
return ScopedFileUsageLimit(scope="user", scope_id=user_api_key_dict.user_id, limit=limit)
return None
def resolve_scoped_limits(
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
setting: FileUsageSetting,
) -> tuple[ScopedFileUsageLimit, ...]:
team_limit: Final = _team_limit(user_api_key_dict, setting)
candidates: Final = (
_caller_counter(user_api_key_dict, _key_limit(user_api_key_dict, general_settings, setting)),
ScopedFileUsageLimit(scope="team", scope_id=user_api_key_dict.team_id, limit=team_limit)
if team_limit is not None and user_api_key_dict.team_id
else None,
)
return tuple(scoped for scoped in candidates if scoped is not None)
async def _increment_all_or_none(
cache: "InternalUsageCache",
counters: tuple[tuple[str, ScopedFileUsageLimit], ...],
ttl_seconds: int,
) -> ScopedFileUsageLimit | None:
if not counters:
return None
(counter_key, scoped), rest = counters[0], counters[1:]
count: Final = await cache.async_increment_cache(
key=counter_key, value=1, litellm_parent_otel_span=None, ttl=ttl_seconds
)
over_here: Final = count is not None and count > scoped.limit.value
exceeded: Final = scoped if over_here else await _increment_all_or_none(cache, rest, ttl_seconds)
if exceeded is not None:
await cache.async_increment_cache(key=counter_key, value=-1, litellm_parent_otel_span=None, ttl=ttl_seconds)
return exceeded
@with_service_target("rate_limits")
async def consume_file_usage(
cache: "InternalUsageCache",
limits: tuple[ScopedFileUsageLimit, ...],
window_seconds: int,
subject: str,
now: float,
) -> FileUsageLimitExceeded | None:
window_start: Final = int(now // window_seconds) * window_seconds
counters: Final = tuple(
(
f"{_COUNTER_PREFIX}:{scoped.limit.setting}:{scoped.scope}:{scoped.scope_id}:{subject}:{window_start}",
scoped,
)
for scoped in limits
)
exceeded: Final = await _increment_all_or_none(cache, counters, window_seconds)
if exceeded is None:
return None
return FileUsageLimitExceeded(
limit=exceeded,
retry_after_seconds=max(1, math.ceil(window_start + window_seconds - now)),
)
def describe_limit_source(source: LimitSource) -> str:
match source:
case "key":
return "in this key's metadata"
case "team":
return "in this team's metadata"
case "general_settings":
return "in general_settings"
return assert_never(source)
def _describe_scope(scoped: ScopedFileUsageLimit) -> str:
match scoped.scope:
case "key":
return "this key"
case "user":
return f"user {scoped.scope_id}"
case "team":
return f"team {scoped.scope_id}"
return assert_never(scoped.scope)
def _raise_limit_exceeded(exceeded: FileUsageLimitExceeded, what_ran_out: str, when_it_resets: str) -> NoReturn:
scoped: Final = exceeded.limit
headers: Final = {"retry-after": str(exceeded.retry_after_seconds)}
raise ProxyException(
message=(
f"{what_ran_out}: {scoped.limit.setting} is {scoped.limit.value} for {_describe_scope(scoped)} "
f"(set {describe_limit_source(scoped.limit.source)}). {when_it_resets}"
),
type="rate_limit_error",
param=None,
code=429,
headers=headers,
)
async def enforce_batch_file_upload_limit(
cache: "InternalUsageCache",
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
clock: Callable[[], float] = time.time,
) -> None:
limits: Final = resolve_scoped_limits(user_api_key_dict, general_settings, MAX_BATCH_FILE_UPLOADS_PER_DAY_KEY)
if not limits:
return
exceeded: Final = await consume_file_usage(cache, limits, _DAY_SECONDS, "", clock())
if exceeded is None:
return
_raise_limit_exceeded(
exceeded,
"Batch file upload limit reached, the file was not forwarded to the provider",
f"The count resets at 00:00 UTC, in {exceeded.retry_after_seconds} seconds.",
)
async def enforce_file_download_limit(
cache: "InternalUsageCache",
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
file_id: str,
clock: Callable[[], float] = time.time,
) -> None:
limits: Final = resolve_scoped_limits(user_api_key_dict, general_settings, MAX_FILE_DOWNLOADS_PER_MINUTE_KEY)
if not limits:
return
exceeded: Final = await consume_file_usage(cache, limits, _MINUTE_SECONDS, file_id, clock())
if exceeded is None:
return
_raise_limit_exceeded(
exceeded,
f"Download limit reached for file {file_id}",
f"Retry in {exceeded.retry_after_seconds} seconds.",
)

View file

@ -84,6 +84,11 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
validate_managed_files_requirement,
validate_managed_id_requirement,
)
from litellm.proxy.openai_files_endpoints.file_usage_caps import (
batch_file_record_limit,
enforce_batch_file_upload_limit,
enforce_file_download_limit,
)
from litellm.proxy.openai_files_endpoints.general_upload_validation import (
MB,
check_allowed_extension,
@ -574,6 +579,7 @@ async def create_file(
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
general_settings_view,
llm_router,
proxy_config,
proxy_logging_obj,
@ -672,6 +678,7 @@ async def create_file(
file_source,
_MAX_BATCH_FILE_SIZE_MB_ADAPTER.validate_python(general_settings.get("max_batch_file_size_mb")),
PASSTHROUGH_BATCH_LINE_SHAPE if passthrough else BATCH_LINE_SHAPE,
batch_file_record_limit(user_api_key_dict, general_settings_view()),
)
if batch_file_failure is not None:
raise_batch_file_validation_failure(batch_file_failure)
@ -748,6 +755,11 @@ async def create_file(
seconds=expires_after_seconds,
)
if purpose == "batch":
await enforce_batch_file_upload_limit(
proxy_logging_obj.file_usage_cache, user_api_key_dict, general_settings_view()
)
# Include original request and headers in the data
data = await add_litellm_data_to_request(
data=data,
@ -964,6 +976,7 @@ async def get_file_content(
"""
from litellm.proxy.proxy_server import (
general_settings,
general_settings_view,
llm_router,
proxy_config,
proxy_logging_obj,
@ -978,6 +991,9 @@ async def get_file_content(
user_api_key_dict=user_api_key_dict,
managed_files_obj=proxy_logging_obj.get_proxy_hook("managed_files"),
)
await enforce_file_download_limit(
proxy_logging_obj.file_usage_cache, user_api_key_dict, general_settings_view(), file_id
)
# Include original request and headers in the data
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
@ -1216,6 +1232,8 @@ async def get_file_content(
)
verbose_proxy_logger.exception("litellm.proxy.proxy_server.retrieve_file_content(): Exception occured - %s", e)
verbose_proxy_logger.debug(traceback.format_exc())
if isinstance(e, ProxyException):
raise e
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),

View file

@ -18293,6 +18293,9 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"admission_queue_timeout_seconds": "Float",
"max_request_size_mb": "Integer",
"max_batch_file_size_mb": "Integer",
"max_batch_file_records": "Integer",
"max_batch_file_uploads_per_day": "Integer",
"max_file_downloads_per_minute": "Integer",
"max_file_size_mb": "Integer",
"allowed_file_extensions": "List",
"blocked_file_extensions": "List",

View file

@ -52,6 +52,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.constants import (
DEFAULT_MODEL_CREATED_AT_TIME,
FILE_USAGE_MAX_TRACKED_COUNTERS,
LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL,
MAX_TEAM_LIST_LIMIT,
PROXY_REJECTED_BEFORE_ROUTING_KEY,
@ -125,6 +126,7 @@ from litellm._logging import _redact_string, verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.exceptions import (
GuardrailRaisedException,
RejectedRequestError,
@ -1213,6 +1215,9 @@ class ProxyLogging:
self.internal_usage_cache: InternalUsageCache = InternalUsageCache(
dual_cache=DualCache(default_in_memory_ttl=1) # ping redis cache every 1s
)
self.file_usage_cache: Final = InternalUsageCache(
dual_cache=DualCache(in_memory_cache=InMemoryCache(max_size_in_memory=FILE_USAGE_MAX_TRACKED_COUNTERS))
)
self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache)
self.cache_control_check = _PROXY_CacheControlCheck()
self.alerting: list[str] | None = None
@ -1354,6 +1359,7 @@ class ProxyLogging:
if redis_cache is not None:
self.internal_usage_cache.dual_cache.redis_cache = redis_cache
self.file_usage_cache.dual_cache.redis_cache = redis_cache
self.db_spend_update_writer.redis_update_buffer.redis_cache = redis_cache
self.db_spend_update_writer.pod_lock_manager.redis_cache = redis_cache

View file

@ -0,0 +1,353 @@
import json
import math
import re
import time
import uuid
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from types import MappingProxyType
from typing import Final
from urllib.parse import unquote, urlsplit
import httpx
import jwt
import yaml
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway, eventually, object_value
from integration._support.wire import Reply, Request
from pydantic import JsonValue
from redis import Redis
RECORDS: Final = "max_batch_file_records"
UPLOADS: Final = "max_batch_file_uploads_per_day"
DOWNLOADS: Final = "max_file_downloads_per_minute"
DAY_SECONDS: Final = 24 * 60 * 60
MINUTE_SECONDS: Final = 60
PROVIDER_KEY: Final = "integration-provider-key"
ROUTED_MODEL: Final = "batch-file-caps-routed"
PROVIDER_REJECTS: Final = "provider-rejects-this-upload"
MISSING_FILE: Final = "file-missing-"
IN_KEY: Final = "in this key's metadata"
IN_TEAM: Final = "in this team's metadata"
IN_GENERAL_SETTINGS: Final = "in general_settings"
THIS_KEY: Final = "this key"
JWT_KEY_ID: Final = "batch-file-caps-signing-key"
COMPLETION_TEXT: Final = "batch file caps completion"
_MARKER: Final = re.compile(rb"caps[0-9a-f]{32}")
_PURPOSE: Final = re.compile(rb'name="purpose"\r\n\r\n([A-Za-z_-]+)')
_CONTENT_PATH: Final = re.compile(r"/v1/files/(.+)/content")
_FILE_PATH: Final = re.compile(r"/v1/files/(.+)")
def marker() -> str:
return "caps" + uuid.uuid4().hex
def batch_line(mark: str, index: int) -> bytes:
return json.dumps(
{
"custom_id": f"{mark}-{index}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "batch file caps"}]},
}
).encode()
def batch_file(mark: str, records: int, separator: bytes = b"\n", ending: bytes = b"\n") -> bytes:
return separator.join(batch_line(mark, index) for index in range(records)) + ending
def file_content(file_id: str) -> bytes:
return json.dumps({"id": "batch_req_1", "custom_id": file_id, "response": {"status_code": 200}}).encode() + b"\n"
def file_object(file_id: str, purpose: str) -> dict[str, JsonValue]:
return {
"id": file_id,
"object": "file",
"bytes": 128,
"created_at": 1700000000,
"filename": "batch.jsonl",
"purpose": purpose,
"status": "processed",
}
def _provider_error(status: int, message: str) -> Reply:
return Reply(
status=status,
body=json.dumps(
{"error": {"message": message, "type": "invalid_request_error", "param": None, "code": None}}
).encode(),
)
def _uploaded(request: Request) -> Reply:
if PROVIDER_REJECTS.encode() in request.body:
return _provider_error(400, "The provider rejected this batch file.")
mark: Final = _MARKER.search(request.body)
purpose: Final = _PURPOSE.search(request.body)
return Reply(
body=json.dumps(
file_object(
f"file-{mark.group().decode()}" if mark is not None else f"file-{uuid.uuid4().hex}",
purpose.group(1).decode() if purpose is not None else "batch",
)
).encode()
)
def _downloaded(file_id: str) -> Reply:
if file_id.startswith(MISSING_FILE):
return _provider_error(404, f"No such File object: {file_id}")
return Reply(body=file_content(file_id), content_type="application/octet-stream")
def _completion() -> Reply:
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": COMPLETION_TEXT},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 7},
}
).encode()
)
def provider(request: Request) -> Reply:
path: Final = urlsplit(request.target).path
if request.method == "POST" and path == "/v1/files":
return _uploaded(request)
if request.method == "POST" and path == "/v1/chat/completions":
return _completion()
content: Final = _CONTENT_PATH.fullmatch(path)
if request.method == "GET" and content is not None:
return _downloaded(unquote(content.group(1)))
described: Final = _FILE_PATH.fullmatch(path)
if request.method == "GET" and described is not None:
return Reply(body=json.dumps(file_object(unquote(described.group(1)), "batch")).encode())
return _provider_error(404, f"No scripted reply for {request.method} {path}")
def seen(requests: tuple[Request, ...], mark: str) -> tuple[Request, ...]:
return tuple(request for request in requests if mark in request.target or mark.encode() in request.body)
def uploads_seen(requests: tuple[Request, ...], mark: str) -> tuple[Request, ...]:
return tuple(
request for request in seen(requests, mark) if (request.method, request.target) == ("POST", "/v1/files")
)
def downloads_seen(requests: tuple[Request, ...], file_id: str) -> tuple[Request, ...]:
return tuple(
request for request in requests if (request.method, request.target) == ("GET", f"/v1/files/{file_id}/content")
)
def assert_forwarded_upload(request: Request, content: bytes) -> None:
assert request.headers["authorization"] == f"Bearer {PROVIDER_KEY}", request.headers
assert content in request.body, request.body
assert b'name="purpose"\r\n\r\nbatch' in request.body, request.body
AUTH_CACHE_TTL_SECONDS: Final = 5
def caps_config(directory: Path, provider_url: str, general_settings: Mapping[str, JsonValue]) -> Path:
base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = directory / f"batch_file_caps_{uuid.uuid4().hex}.yaml"
path.write_text(
yaml.safe_dump(
{
**base,
"model_list": [
{
"model_name": ROUTED_MODEL,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": PROVIDER_KEY,
"api_base": f"{provider_url}/v1",
},
}
],
"general_settings": {
**base["general_settings"],
"user_api_key_cache_ttl": AUTH_CACHE_TTL_SECONDS,
**general_settings,
},
"files_settings": [
{"custom_llm_provider": "openai", "api_base": f"{provider_url}/v1", "api_key": PROVIDER_KEY}
],
}
)
)
return path
def provider_environment(provider_url: str) -> dict[str, str]:
return {"OPENAI_BASE_URL": f"{provider_url}/v1", "OPENAI_API_KEY": PROVIDER_KEY}
def upload(
candidate: Gateway,
key: str,
content: bytes,
*,
path: str = "/v1/files",
purpose: str = "batch",
filename: str = "batch.jsonl",
fields: Mapping[str, str] = MappingProxyType({}),
) -> httpx.Response:
return candidate.request_multipart(
path, {"purpose": purpose, **fields}, {"file": (filename, content, "application/jsonl")}, key=key
)
def download(
candidate: Gateway,
key: str,
file_id: str,
*,
route: str = "/v1/files/{}/content",
params: Mapping[str, str] | None = None,
) -> httpx.Response:
return candidate.request("GET", route.format(file_id), key=key, params=params)
def window_end(window_seconds: int, room_seconds: int) -> float:
started: Final = eventually(
time.time,
lambda now: window_seconds - now % window_seconds >= room_seconds,
seconds=room_seconds + 5,
)
return (started // window_seconds + 1) * window_seconds
@dataclass(frozen=True, slots=True)
class Timed:
response: httpx.Response
before: float
after: float
def timed(send: Callable[[], httpx.Response]) -> Timed:
before: Final = time.time()
response: Final = send()
return Timed(response, before, time.time())
def _assert_rate_limited(observed: Timed, ends: float, what: str, held: str, reset: str) -> None:
response: Final = observed.response
assert response.status_code == 429, response.text
assert observed.after < ends, f"The counted sequence ran past its window: {observed.after} >= {ends}"
retry_after: Final = int(response.headers["retry-after"])
earliest: Final = max(1, math.ceil(ends - observed.after))
latest: Final = max(1, math.ceil(ends - observed.before))
assert earliest <= retry_after <= latest, (retry_after, earliest, latest)
assert response.json() == {
"error": {
"message": f"{what}: {held}. {reset.format(retry_after)}",
"type": "rate_limit_error",
"param": None,
"code": "429",
}
}, response.text
def assert_upload_limited(observed: Timed, day_ends: float, limit: int, holder: str, source: str) -> None:
_assert_rate_limited(
observed,
day_ends,
"Batch file upload limit reached, the file was not forwarded to the provider",
f"{UPLOADS} is {limit} for {holder} (set {source})",
"The count resets at 00:00 UTC, in {} seconds.",
)
def assert_download_limited(
observed: Timed, minute_ends: float, file_id: str, limit: int, holder: str, source: str
) -> None:
_assert_rate_limited(
observed,
minute_ends,
f"Download limit reached for file {file_id}",
f"{DOWNLOADS} is {limit} for {holder} (set {source})",
"Retry in {} seconds.",
)
def assert_too_many_records(response: httpx.Response, limit: int, source: str) -> None:
assert response.status_code == 413, response.text
assert response.json() == {
"error": {
"message": (
f"Batch input file has more than {limit} records, which exceeds the {RECORDS} of {limit} "
f"set {source}. The file was not forwarded to the provider."
),
"type": "invalid_request_error",
"param": "file",
"code": "413",
}
}, response.text
@dataclass(frozen=True, slots=True)
class Counter:
name: str
count: int
ttl: int
def hashed(key: str) -> str:
return sha256(key.encode()).hexdigest()
def counters(cache: Redis, setting: str, holder: str) -> tuple[Counter, ...]:
return tuple(
Counter(name.decode(), int(cache.get(name) or 0), cache.ttl(name))
for name in sorted(cache.scan_iter(match=f"*litellm:file_usage:{setting}:*{holder}*", count=1000))
)
def config_entry(gateway: Gateway, setting: str) -> Mapping[str, JsonValue]:
response: Final = gateway.request("GET", "/config/list", params={"config_type": "general_settings"})
assert response.status_code == 200, response.text
(entry,) = (object_value(field) for field in response.json() if field["field_name"] == setting)
return entry
@dataclass(frozen=True, slots=True)
class Signer:
private_key: rsa.RSAPrivateKey
jwks: bytes
def signer() -> Signer:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())
return Signer(private_key, json.dumps({"keys": [{**json.loads(public_jwk), "kid": JWT_KEY_ID}]}).encode())
def signed_token(identity: Signer, subject: str) -> str:
now: Final = int(time.time())
return jwt.encode(
{"sub": subject, "iat": now, "exp": now + 300, "jti": uuid.uuid4().hex},
identity.private_key,
algorithm="RS256",
headers={"kid": JWT_KEY_ID},
)

View file

@ -0,0 +1,928 @@
import asyncio
import os
import time
import uuid
from collections.abc import AsyncIterator, Iterator, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from functools import partial
from types import MappingProxyType
from typing import Final
from urllib.parse import unquote
import httpx
import openai
import pytest
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, string_value
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.management._batch_file_caps import (
COMPLETION_TEXT,
DAY_SECONDS,
DOWNLOADS,
IN_GENERAL_SETTINGS,
IN_KEY,
IN_TEAM,
MINUTE_SECONDS,
MISSING_FILE,
PROVIDER_KEY,
PROVIDER_REJECTS,
RECORDS,
ROUTED_MODEL,
THIS_KEY,
UPLOADS,
Signer,
Timed,
assert_download_limited,
assert_forwarded_upload,
assert_too_many_records,
assert_upload_limited,
batch_file,
batch_line,
caps_config,
counters,
download,
downloads_seen,
file_content,
hashed,
marker,
provider,
provider_environment,
seen,
signed_token,
signer,
timed,
upload,
uploads_seen,
window_end,
)
from pydantic import JsonValue
from redis import Redis
pytestmark = pytest.mark.timeout(600)
YAML_RECORDS: Final = 3
YAML_UPLOADS: Final = 4
ROOMY: Final = 1000
DAY_ROOM_SECONDS: Final = 120
MINUTE_ROOM_SECONDS: Final = 30
INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation"
UPLOAD_ROUTES: Final = ("/v1/files", "/files", "/openai/v1/files")
DOWNLOAD_ROUTES: Final = ("/v1/files/{}/content", "/files/{}/content", "/openai/v1/files/{}/content")
ROTATIONS: Final = (
pytest.param(0, id="openai-prefixed-route-last"),
pytest.param(1, id="v1-route-last"),
pytest.param(2, id="bare-route-last"),
)
LEVELS: Final = ("key", "team")
MANAGED: Final = MappingProxyType({"target_model_names": ROUTED_MODEL})
HOSTILE_IDS: Final = (
pytest.param("f" * 5120, id="5kb-id"),
pytest.param("file-{}%20x", id="percent-encoded-space"),
)
SLASHED_IDS: Final = (
pytest.param("file-{}/x", id="slash"),
pytest.param("file-{}%2Fx", id="percent-encoded-slash"),
)
@dataclass(frozen=True, slots=True)
class Rig:
candidate: Gateway
sibling: Gateway
provider: Wire
cache: Redis
identity: Signer
def gateways(self, count: int) -> tuple[Gateway, ...]:
return tuple((self.candidate, self.sibling)[index % 2] for index in range(count))
def _subscribers(cache: Redis) -> int:
return int(cache.pubsub_numsub(INVALIDATION_CHANNEL)[0][1])
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
directory: Final = tmp_path_factory.mktemp("batch_file_caps")
identity: Final = signer()
def jwks(_request: Request) -> Reply:
return Reply(body=identity.jwks)
with (
gateway_from_environment() as gateway,
wire_server(provider) as files,
wire_server(jwks) as keys,
Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache,
):
config: Final = caps_config(
directory,
files.url,
{
RECORDS: YAML_RECORDS,
UPLOADS: YAML_UPLOADS,
"enable_jwt_auth": True,
"litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True},
},
)
environment: Final = {**provider_environment(files.url), "JWT_PUBLIC_KEY_URL": keys.url}
subscribed: Final = _subscribers(cache)
with (
owned_proxy(gateway, directory, environment, config=config, workers=2) as candidate,
owned_proxy(gateway, directory, environment, config=config) as sibling,
):
eventually(partial(_subscribers, cache), lambda count: count >= subscribed + 3, seconds=60)
yield Rig(candidate, sibling, files, cache, identity)
def _key(scenario: Scenario, team: str | None = None, **limits: JsonValue) -> str:
return scenario.key(metadata=limits, **({"team_id": team} if team is not None else {}))
def _accepted(response: httpx.Response) -> str:
assert response.status_code == 200, response.text
return string_value(response.json()["id"])
def _statuses(responses: tuple[httpx.Response, ...]) -> list[int]:
return [response.status_code for response in responses]
def _counts(rig: Rig, setting: str, holder: str) -> list[int]:
return [counter.count for counter in counters(rig.cache, setting, holder)]
def _rotated(routes: tuple[str, str, str], first: int) -> tuple[str, ...]:
return routes[first:] + routes[:first]
def test_yaml_record_limit_rejects_a_longer_batch_file_before_it_reaches_the_provider(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = scenario.key()
over: Final = marker()
within: Final = marker()
content: Final = batch_file(within, YAML_RECORDS)
assert_too_many_records(
upload(rig.candidate, key, batch_file(over, YAML_RECORDS + 1)), YAML_RECORDS, IN_GENERAL_SETTINGS
)
assert _accepted(upload(rig.sibling, key, content)) == f"file-{within}"
requests: Final = rig.provider.drain()
assert seen(requests, over) == ()
(forwarded,) = uploads_seen(requests, within)
assert_forwarded_upload(forwarded, content)
@pytest.mark.parametrize(
("key_limit", "team_limit", "limit", "source"),
[
pytest.param(2, None, 2, IN_KEY, id="key-below-yaml"),
pytest.param(5, None, 5, IN_KEY, id="key-above-yaml"),
pytest.param(None, 2, 2, IN_TEAM, id="team-below-yaml"),
pytest.param(5, 4, 4, IN_TEAM, id="team-below-key"),
pytest.param(2, 4, 2, IN_KEY, id="key-below-team"),
pytest.param(None, 5, YAML_RECORDS, IN_GENERAL_SETTINGS, id="yaml-below-team"),
],
)
def test_record_limit_is_the_lower_of_the_key_level_and_team_limits(
rig: Rig, key_limit: int | None, team_limit: int | None, limit: int, source: str
) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata={RECORDS: team_limit}) if team_limit is not None else None
key: Final = _key(scenario, team, **({RECORDS: key_limit} if key_limit is not None else {}))
over: Final = marker()
within: Final = marker()
content: Final = batch_file(within, limit)
assert_too_many_records(upload(rig.candidate, key, batch_file(over, limit + 1)), limit, source)
assert _accepted(upload(rig.sibling, key, content)) == f"file-{within}"
requests: Final = rig.provider.drain()
assert seen(requests, over) == ()
(forwarded,) = uploads_seen(requests, within)
assert_forwarded_upload(forwarded, content)
@pytest.mark.parametrize(
("separator", "ending", "records"),
[
pytest.param(b"\n\n\n", b"\n\n", YAML_RECORDS, id="blank-lines-at-the-limit"),
pytest.param(b"\n \n", b"\n\t\n", YAML_RECORDS + 1, id="blank-lines-over-the-limit"),
pytest.param(b"\r\n", b"", YAML_RECORDS, id="crlf-without-trailing-newline-at-the-limit"),
pytest.param(b"\r\n", b"", YAML_RECORDS + 1, id="crlf-without-trailing-newline-over-the-limit"),
],
)
def test_record_limit_counts_request_lines_not_blank_lines_or_line_endings(
rig: Rig, separator: bytes, ending: bytes, records: int
) -> None:
with rig.candidate.scenario() as scenario:
key: Final = scenario.key()
mark: Final = marker()
content: Final = batch_file(mark, records, separator, ending)
response: Final = upload(rig.candidate, key, content)
requests: Final = uploads_seen(rig.provider.drain(), mark)
if records > YAML_RECORDS:
assert_too_many_records(response, YAML_RECORDS, IN_GENERAL_SETTINGS)
assert requests == ()
return
assert _accepted(response) == f"file-{mark}"
(forwarded,) = requests
assert_forwarded_upload(forwarded, content)
def test_only_batch_purpose_uploads_are_limited_and_counted(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{RECORDS: 1, UPLOADS: 1})
mark: Final = marker()
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
purposes: Final = ("user_data", "fine-tune", "user_data", "fine-tune")
others: Final = tuple(
upload(gateway, key, batch_file(mark, 5), purpose=purpose, filename="notes.jsonl")
for gateway, purpose in zip(rig.gateways(4), purposes, strict=True)
)
assert _statuses(others) == [200] * 4, [response.text for response in others]
assert _accepted(upload(rig.candidate, key, batch_file(mark, 1))) == f"file-{mark}"
assert_upload_limited(
timed(partial(upload, rig.sibling, key, batch_file(mark, 1))), day_ends, 1, THIS_KEY, IN_KEY
)
forwarded: Final = uploads_seen(rig.provider.drain(), mark)
assert [request.body.count(b"\r\n\r\nbatch\r\n") for request in forwarded] == [0, 0, 0, 0, 1], forwarded
assert _counts(rig, UPLOADS, hashed(key)) == [1]
def test_yaml_daily_upload_limit_counts_one_key_across_proxy_processes(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = scenario.key()
other_key: Final = scenario.key()
mark: Final = marker()
other_mark: Final = marker()
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
accepted: Final = tuple(upload(gateway, key, content) for gateway in rig.gateways(YAML_UPLOADS))
assert _statuses(accepted) == [200] * YAML_UPLOADS, [response.text for response in accepted]
for gateway in rig.gateways(2):
assert_upload_limited(
timed(partial(upload, gateway, key, content)), day_ends, YAML_UPLOADS, THIS_KEY, IN_GENERAL_SETTINGS
)
assert _accepted(upload(rig.candidate, other_key, batch_file(other_mark, 1))) == f"file-{other_mark}"
requests: Final = rig.provider.drain()
forwarded: Final = uploads_seen(requests, mark)
assert len(forwarded) == YAML_UPLOADS, forwarded
for request in forwarded:
assert_forwarded_upload(request, content)
assert len(uploads_seen(requests, other_mark)) == 1
(counter,) = counters(rig.cache, UPLOADS, hashed(key))
assert counter.count == YAML_UPLOADS, counter
assert 0 < counter.ttl <= DAY_SECONDS, counter
assert counter.name.endswith(
f"litellm:file_usage:{UPLOADS}:key:{hashed(key)}::{int(day_ends) - DAY_SECONDS}"
), counter
@pytest.mark.parametrize("limit", [2, 6])
def test_key_daily_upload_limit_replaces_the_yaml_limit(rig: Rig, limit: int) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: limit})
mark: Final = marker()
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
accepted: Final = tuple(upload(gateway, key, content) for gateway in rig.gateways(limit))
assert _statuses(accepted) == [200] * limit, [response.text for response in accepted]
assert_upload_limited(timed(partial(upload, rig.candidate, key, content)), day_ends, limit, THIS_KEY, IN_KEY)
assert len(uploads_seen(rig.provider.drain(), mark)) == limit
assert _counts(rig, UPLOADS, hashed(key)) == [limit]
def test_upload_rejected_by_the_key_limit_does_not_use_a_team_slot(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata={UPLOADS: 5})
tight_key: Final = _key(scenario, team, **{UPLOADS: 2})
loose_key: Final = _key(scenario, team)
tight_mark: Final = marker()
loose_mark: Final = marker()
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
tight: Final = tuple(upload(gateway, tight_key, batch_file(tight_mark, 1)) for gateway in rig.gateways(2))
assert _statuses(tight) == [200, 200], [response.text for response in tight]
for gateway in rig.gateways(2):
assert_upload_limited(
timed(partial(upload, gateway, tight_key, batch_file(tight_mark, 1))), day_ends, 2, THIS_KEY, IN_KEY
)
loose: Final = tuple(upload(gateway, loose_key, batch_file(loose_mark, 1)) for gateway in rig.gateways(3))
assert _statuses(loose) == [200, 200, 200], [response.text for response in loose]
assert_upload_limited(
timed(partial(upload, rig.sibling, loose_key, batch_file(loose_mark, 1))),
day_ends,
5,
f"team {team}",
IN_TEAM,
)
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, tight_mark)) == 2
assert len(uploads_seen(requests, loose_mark)) == 3
assert _counts(rig, UPLOADS, hashed(tight_key)) == [2]
assert _counts(rig, UPLOADS, hashed(loose_key)) == [3]
assert _counts(rig, UPLOADS, f"team:{team}:") == [5]
def test_rejected_uploads_do_not_use_a_daily_upload_slot(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = scenario.key()
rejected_mark: Final = marker()
mark: Final = marker()
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
assert_too_many_records(
upload(rig.candidate, key, batch_file(rejected_mark, YAML_RECORDS + 1)), YAML_RECORDS, IN_GENERAL_SETTINGS
)
bad_expiry: Final = upload(
rig.sibling,
key,
batch_file(rejected_mark, 1),
fields={"expires_after[anchor]": "created_at", "expires_after[seconds]": "soon"},
)
assert bad_expiry.status_code == 400, bad_expiry.text
assert "expires_after[seconds] must be a valid integer, got 'soon'" in bad_expiry.text
invalid_line: Final = upload(rig.candidate, key, batch_line(rejected_mark, 0) + b"\n{not json\n")
assert invalid_line.status_code == 400, invalid_line.text
assert "Batch input file line 2 is not valid JSON" in invalid_line.text
wrong_extension: Final = upload(rig.sibling, key, batch_file(rejected_mark, 1), filename="batch.txt")
assert wrong_extension.status_code == 400, wrong_extension.text
assert "Batch input files must be .jsonl files" in wrong_extension.text
accepted: Final = tuple(upload(gateway, key, content) for gateway in rig.gateways(YAML_UPLOADS))
assert _statuses(accepted) == [200] * YAML_UPLOADS, [response.text for response in accepted]
assert_upload_limited(
timed(partial(upload, rig.candidate, key, content)), day_ends, YAML_UPLOADS, THIS_KEY, IN_GENERAL_SETTINGS
)
requests: Final = rig.provider.drain()
assert seen(requests, rejected_mark) == ()
assert len(uploads_seen(requests, mark)) == YAML_UPLOADS
assert _counts(rig, UPLOADS, hashed(key)) == [YAML_UPLOADS]
def test_callers_without_a_valid_key_are_rejected_before_anything_is_counted(rig: Rig) -> None:
mark: Final = marker()
unknown_key: Final = f"sk-{mark}"
file_id: Final = f"file-{mark}"
responses: Final = (
rig.candidate.client.post(
"/v1/files",
data={"purpose": "batch"},
files={"file": ("batch.jsonl", batch_file(mark, 1), "application/jsonl")},
),
upload(rig.sibling, unknown_key, batch_file(mark, 1)),
rig.candidate.client.get(f"/v1/files/{file_id}/content"),
download(rig.sibling, unknown_key, file_id),
)
assert _statuses(responses) == [401, 401, 401, 401], [response.text for response in responses]
assert seen(rig.provider.drain(), mark) == ()
assert counters(rig.cache, UPLOADS, hashed(unknown_key)) == ()
assert counters(rig.cache, DOWNLOADS, mark) == ()
def test_upload_the_provider_rejects_still_uses_a_daily_upload_slot(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 2})
mark: Final = marker()
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
refused: Final = upload(rig.candidate, key, batch_file(f"{mark}-{PROVIDER_REJECTS}", 1))
assert refused.status_code == 400, refused.text
assert "The provider rejected this batch file." in refused.text
assert _accepted(upload(rig.sibling, key, content)) == f"file-{mark}"
assert_upload_limited(timed(partial(upload, rig.candidate, key, content)), day_ends, 2, THIS_KEY, IN_KEY)
assert len(uploads_seen(rig.provider.drain(), mark)) == 2
assert _counts(rig, UPLOADS, hashed(key)) == [2]
def test_openai_sdk_caller_sees_each_limit_as_a_status_error(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 2, DOWNLOADS: 2})
mark: Final = marker()
file_id: Final = f"file-{mark}"
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
with (
httpx.Client(timeout=15, trust_env=False) as transport,
openai.OpenAI(
base_url=f"{str(rig.candidate.client.base_url).rstrip('/')}/v1",
api_key=key,
max_retries=0,
http_client=transport,
) as client,
):
with pytest.raises(openai.APIStatusError) as too_long:
client.files.create(file=("batch.jsonl", batch_file(marker(), YAML_RECORDS + 1)), purpose="batch")
assert_too_many_records(too_long.value.response, YAML_RECORDS, IN_GENERAL_SETTINGS)
created: Final = tuple(
client.files.create(file=("batch.jsonl", content), purpose="batch") for _ in range(2)
)
assert [file.id for file in created] == [file_id, file_id]
upload_started: Final = time.time()
with pytest.raises(openai.RateLimitError) as upload_limited:
client.files.create(file=("batch.jsonl", content), purpose="batch")
assert_upload_limited(
Timed(upload_limited.value.response, upload_started, time.time()), day_ends, 2, THIS_KEY, IN_KEY
)
downloaded: Final = tuple(client.files.content(file_id).content for _ in range(2))
assert downloaded == (file_content(file_id), file_content(file_id))
download_started: Final = time.time()
with pytest.raises(openai.RateLimitError) as download_limited:
client.files.content(file_id)
assert_download_limited(
Timed(download_limited.value.response, download_started, time.time()),
minute_ends,
file_id,
2,
THIS_KEY,
IN_KEY,
)
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == 2
assert len(downloads_seen(requests, file_id)) == 2
async def test_async_openai_sdk_caller_sees_each_limit_as_a_status_error(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 2, DOWNLOADS: 2})
mark: Final = marker()
file_id: Final = f"file-{mark}"
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
async with (
httpx.AsyncClient(timeout=15, trust_env=False) as transport,
openai.AsyncOpenAI(
base_url=f"{str(rig.sibling.client.base_url).rstrip('/')}/v1",
api_key=key,
max_retries=0,
http_client=transport,
) as client,
):
with pytest.raises(openai.APIStatusError) as too_long:
await client.files.create(file=("batch.jsonl", batch_file(marker(), YAML_RECORDS + 1)), purpose="batch")
assert_too_many_records(too_long.value.response, YAML_RECORDS, IN_GENERAL_SETTINGS)
first: Final = await client.files.create(file=("batch.jsonl", content), purpose="batch")
second: Final = await client.files.create(file=("batch.jsonl", content), purpose="batch")
assert [first.id, second.id] == [file_id, file_id]
upload_started: Final = time.time()
with pytest.raises(openai.RateLimitError) as upload_limited:
await client.files.create(file=("batch.jsonl", content), purpose="batch")
assert_upload_limited(
Timed(upload_limited.value.response, upload_started, time.time()), day_ends, 2, THIS_KEY, IN_KEY
)
first_download: Final = await client.files.content(file_id)
second_download: Final = await client.files.content(file_id)
assert [first_download.content, second_download.content] == [file_content(file_id), file_content(file_id)]
download_started: Final = time.time()
with pytest.raises(openai.RateLimitError) as download_limited:
await client.files.content(file_id)
assert_download_limited(
Timed(download_limited.value.response, download_started, time.time()),
minute_ends,
file_id,
2,
THIS_KEY,
IN_KEY,
)
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == 2
assert len(downloads_seen(requests, file_id)) == 2
@pytest.mark.parametrize("first", ROTATIONS)
def test_every_upload_route_shares_one_daily_count(rig: Rig, first: int) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 2})
mark: Final = marker()
content: Final = batch_file(mark, 1)
allowed_first, allowed_second, limited = _rotated(UPLOAD_ROUTES, first)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
assert _accepted(upload(rig.candidate, key, content, path=allowed_first)) == f"file-{mark}"
assert _accepted(upload(rig.sibling, key, content, path=allowed_second)) == f"file-{mark}"
assert_upload_limited(
timed(partial(upload, rig.candidate, key, content, path=limited)), day_ends, 2, THIS_KEY, IN_KEY
)
assert len(uploads_seen(rig.provider.drain(), mark)) == 2
assert _counts(rig, UPLOADS, hashed(key)) == [2]
def test_jwt_callers_are_counted_per_user_across_tokens(rig: Rig) -> None:
subject: Final = f"caps-jwt-{uuid.uuid4().hex}"
other_subject: Final = f"caps-jwt-{uuid.uuid4().hex}"
with rig.candidate.scenario() as scenario:
mark: Final = marker()
other_mark: Final = marker()
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
first: Final = upload(rig.candidate, signed_token(rig.identity, subject), content)
scenario.cleanups.callback(scenario.delete_user, subject)
assert _accepted(first) == f"file-{mark}"
rest: Final = tuple(
upload(gateway, signed_token(rig.identity, subject), content) for gateway in rig.gateways(YAML_UPLOADS - 1)
)
assert _statuses(rest) == [200] * (YAML_UPLOADS - 1), [response.text for response in rest]
assert_upload_limited(
timed(partial(upload, rig.sibling, signed_token(rig.identity, subject), content)),
day_ends,
YAML_UPLOADS,
f"user {subject}",
IN_GENERAL_SETTINGS,
)
other: Final = upload(rig.candidate, signed_token(rig.identity, other_subject), batch_file(other_mark, 1))
scenario.cleanups.callback(scenario.delete_user, other_subject)
assert _accepted(other) == f"file-{other_mark}"
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == YAML_UPLOADS
assert len(uploads_seen(requests, other_mark)) == 1
assert _counts(rig, UPLOADS, f"user:{subject}:") == [YAML_UPLOADS]
def test_key_download_limit_counts_one_file_per_key_across_proxy_processes(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 3})
other_key: Final = _key(scenario, **{DOWNLOADS: 3})
file_id: Final = f"file-{marker()}"
other_file: Final = f"file-{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
allowed: Final = tuple(download(gateway, key, file_id) for gateway in rig.gateways(3))
assert _statuses(allowed) == [200, 200, 200], [response.text for response in allowed]
assert [response.content for response in allowed] == [file_content(file_id)] * 3
for gateway in rig.gateways(2):
assert_download_limited(
timed(partial(download, gateway, key, file_id)), minute_ends, file_id, 3, THIS_KEY, IN_KEY
)
same_key_other_file: Final = download(rig.candidate, key, other_file)
assert (same_key_other_file.status_code, same_key_other_file.content) == (200, file_content(other_file))
other_key_same_file: Final = download(rig.sibling, other_key, file_id)
assert (other_key_same_file.status_code, other_key_same_file.content) == (200, file_content(file_id))
requests: Final = rig.provider.drain()
forwarded: Final = downloads_seen(requests, file_id)
assert len(forwarded) == 4, forwarded
assert {request.headers["authorization"] for request in forwarded} == {f"Bearer {PROVIDER_KEY}"}
assert len(downloads_seen(requests, other_file)) == 1
(counter,) = counters(rig.cache, DOWNLOADS, f"{hashed(key)}:{file_id}:")
assert counter.count == 3, counter
assert 0 < counter.ttl <= MINUTE_SECONDS, counter
assert counter.name.endswith(
f"litellm:file_usage:{DOWNLOADS}:key:{hashed(key)}:{file_id}:{int(minute_ends) - MINUTE_SECONDS}"
), counter
def test_team_download_limit_is_shared_by_every_key_on_the_team(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata={DOWNLOADS: 3})
first_key: Final = _key(scenario, team)
second_key: Final = _key(scenario, team)
outside_key: Final = scenario.key()
file_id: Final = f"file-{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
allowed: Final = (
download(rig.candidate, first_key, file_id),
download(rig.sibling, first_key, file_id),
download(rig.candidate, second_key, file_id),
)
assert _statuses(allowed) == [200, 200, 200], [response.text for response in allowed]
for gateway, key in zip(rig.gateways(2), (second_key, first_key), strict=True):
assert_download_limited(
timed(partial(download, gateway, key, file_id)), minute_ends, file_id, 3, f"team {team}", IN_TEAM
)
outside: Final = download(rig.sibling, outside_key, file_id)
assert (outside.status_code, outside.content) == (200, file_content(file_id)), outside.text
assert len(downloads_seen(rig.provider.drain(), file_id)) == 4
assert _counts(rig, DOWNLOADS, f"team:{team}:{file_id}:") == [3]
def test_downloads_are_unlimited_when_no_limit_is_set(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = scenario.key()
file_id: Final = f"file-{marker()}"
responses: Final = tuple(download(gateway, key, file_id) for gateway in rig.gateways(12))
assert _statuses(responses) == [200] * 12, [response.text for response in responses]
assert {response.content for response in responses} == {file_content(file_id)}
assert len(downloads_seen(rig.provider.drain(), file_id)) == 12
assert counters(rig.cache, DOWNLOADS, file_id) == ()
def test_download_the_provider_cannot_find_still_uses_a_slot(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 2})
file_id: Final = f"{MISSING_FILE}{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
missing: Final = tuple(download(gateway, key, file_id) for gateway in rig.gateways(2))
assert _statuses(missing) == [404, 404], [response.text for response in missing]
assert f"No such File object: {file_id}" in missing[0].text
assert_download_limited(
timed(partial(download, rig.candidate, key, file_id)), minute_ends, file_id, 2, THIS_KEY, IN_KEY
)
assert len(downloads_seen(rig.provider.drain(), file_id)) == 2
@pytest.mark.parametrize("first", ROTATIONS)
def test_every_download_route_shares_one_count_per_file(rig: Rig, first: int) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 2})
file_id: Final = f"file-{marker()}"
allowed_first, allowed_second, limited = _rotated(DOWNLOAD_ROUTES, first)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
allowed: Final = (
download(rig.candidate, key, file_id, route=allowed_first),
download(rig.sibling, key, file_id, route=allowed_second),
)
assert _statuses(allowed) == [200, 200], [response.text for response in allowed]
assert_download_limited(
timed(partial(download, rig.candidate, key, file_id, route=limited)),
minute_ends,
file_id,
2,
THIS_KEY,
IN_KEY,
)
assert len(downloads_seen(rig.provider.drain(), file_id)) == 2
def test_model_routed_files_are_counted_under_the_id_the_caller_uses(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 2})
mark: Final = marker()
provider_id: Final = f"file-{mark}"
plain_id: Final = f"file-{marker()}"
routed_id: Final = _accepted(upload(rig.candidate, key, batch_file(mark, 1), fields={"model": ROUTED_MODEL}))
assert routed_id != provider_id
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
routed: Final = tuple(download(gateway, key, routed_id) for gateway in rig.gateways(2))
assert _statuses(routed) == [200, 200], [response.text for response in routed]
assert {response.content for response in routed} == {file_content(provider_id)}
assert_download_limited(
timed(partial(download, rig.candidate, key, routed_id)), minute_ends, routed_id, 2, THIS_KEY, IN_KEY
)
by_query: Final = tuple(
download(gateway, key, plain_id, params={"model": ROUTED_MODEL}) for gateway in rig.gateways(2)
)
assert _statuses(by_query) == [200, 200], [response.text for response in by_query]
assert_download_limited(
timed(partial(download, rig.sibling, key, plain_id, params={"model": ROUTED_MODEL})),
minute_ends,
plain_id,
2,
THIS_KEY,
IN_KEY,
)
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == 1
assert len(downloads_seen(requests, provider_id)) == 2
assert len(downloads_seen(requests, plain_id)) == 2
def test_managed_files_are_capped_and_counted_under_their_unified_id(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 1, DOWNLOADS: 2})
mark: Final = marker()
provider_id: Final = f"file-{mark}"
rejected_mark: Final = marker()
limited_mark: Final = marker()
assert_too_many_records(
upload(rig.candidate, key, batch_file(rejected_mark, YAML_RECORDS + 1), fields=MANAGED),
YAML_RECORDS,
IN_GENERAL_SETTINGS,
)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
unified_id: Final = _accepted(upload(rig.sibling, key, batch_file(mark, 1), fields=MANAGED))
assert unified_id != provider_id
assert_upload_limited(
timed(partial(upload, rig.candidate, key, batch_file(limited_mark, 1), fields=MANAGED)),
day_ends,
1,
THIS_KEY,
IN_KEY,
)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
served: Final = tuple(download(gateway, key, unified_id) for gateway in rig.gateways(2))
assert _statuses(served) == [200, 200], [response.text for response in served]
assert {response.content for response in served} == {file_content(provider_id)}
assert_download_limited(
timed(partial(download, rig.sibling, key, unified_id)), minute_ends, unified_id, 2, THIS_KEY, IN_KEY
)
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == 1
assert seen(requests, rejected_mark) == ()
assert seen(requests, limited_mark) == ()
assert len(downloads_seen(requests, provider_id)) == 2
assert _counts(rig, UPLOADS, hashed(key)) == [1]
assert _counts(rig, DOWNLOADS, f"{hashed(key)}:{unified_id}:") == [2]
@pytest.mark.parametrize("shape", HOSTILE_IDS)
def test_hostile_file_ids_are_counted_as_the_provider_receives_them(rig: Rig, shape: str) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 2})
sent: Final = shape.format(marker())
counted: Final = unquote(sent)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
served: Final = tuple(download(gateway, key, sent) for gateway in rig.gateways(2))
assert _statuses(served) == [200, 200], [response.text for response in served]
assert {response.content for response in served} == {file_content(counted)}
assert_download_limited(
timed(partial(download, rig.candidate, key, sent)), minute_ends, counted, 2, THIS_KEY, IN_KEY
)
assert len(downloads_seen(rig.provider.drain(), sent)) == 2
assert _counts(rig, DOWNLOADS, f"{hashed(key)}:{counted}:") == [2]
@pytest.mark.parametrize("shape", SLASHED_IDS)
def test_a_file_id_with_a_slash_never_reaches_the_file_route_or_a_counter(rig: Rig, shape: str) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 2})
mark: Final = marker()
answers: Final = tuple(download(gateway, key, shape.format(mark)) for gateway in rig.gateways(3))
assert _statuses(answers) == [401, 401, 401], [response.text for response in answers]
assert seen(rig.provider.drain(), mark) == ()
assert counters(rig.cache, DOWNLOADS, hashed(key)) == ()
def _every_limit(value: JsonValue) -> dict[str, JsonValue]:
return {RECORDS: value, UPLOADS: value, DOWNLOADS: value}
@pytest.mark.parametrize("level", LEVELS)
@pytest.mark.parametrize(
"metadata",
[
pytest.param(_every_limit(0), id="zero"),
pytest.param(_every_limit(-1), id="negative"),
pytest.param(_every_limit(""), id="empty-string"),
pytest.param(_every_limit("abc"), id="word"),
pytest.param(_every_limit([5]), id="list"),
pytest.param(_every_limit({"a": 1}), id="object"),
pytest.param(_every_limit(5.5), id="fraction"),
pytest.param(_every_limit("x" * 5120), id="5kb-string"),
pytest.param(_every_limit(None), id="null"),
pytest.param({}, id="empty-metadata"),
],
)
def test_malformed_limits_are_ignored_and_the_yaml_limits_still_apply(
rig: Rig, metadata: Mapping[str, JsonValue], level: str
) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata=dict(metadata)) if level == "team" else None
key: Final = _key(scenario, team, **(metadata if level == "key" else {}))
over: Final = marker()
mark: Final = marker()
file_id: Final = f"file-{mark}"
content: Final = batch_file(mark, YAML_RECORDS)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
assert_too_many_records(
upload(rig.candidate, key, batch_file(over, YAML_RECORDS + 1)), YAML_RECORDS, IN_GENERAL_SETTINGS
)
accepted: Final = tuple(upload(gateway, key, content) for gateway in rig.gateways(YAML_UPLOADS))
assert _statuses(accepted) == [200] * YAML_UPLOADS, [response.text for response in accepted]
assert_upload_limited(
timed(partial(upload, rig.candidate, key, content)), day_ends, YAML_UPLOADS, THIS_KEY, IN_GENERAL_SETTINGS
)
downloads: Final = tuple(download(gateway, key, file_id) for gateway in rig.gateways(6))
assert _statuses(downloads) == [200] * 6, [response.text for response in downloads]
requests: Final = rig.provider.drain()
assert seen(requests, over) == ()
assert len(uploads_seen(requests, mark)) == YAML_UPLOADS
assert len(downloads_seen(requests, file_id)) == 6
@pytest.mark.parametrize("level", LEVELS)
def test_numeric_string_limits_are_honored(rig: Rig, level: str) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata=_every_limit("2")) if level == "team" else None
key: Final = _key(scenario, team, **(_every_limit("2") if level == "key" else {}))
holder: Final = THIS_KEY if level == "key" else f"team {team}"
source: Final = IN_KEY if level == "key" else IN_TEAM
over: Final = marker()
mark: Final = marker()
file_id: Final = f"file-{mark}"
content: Final = batch_file(mark, 2)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
assert_too_many_records(upload(rig.candidate, key, batch_file(over, 3)), 2, source)
accepted: Final = tuple(upload(gateway, key, content) for gateway in rig.gateways(2))
assert _statuses(accepted) == [200, 200], [response.text for response in accepted]
assert_upload_limited(timed(partial(upload, rig.candidate, key, content)), day_ends, 2, holder, source)
downloads: Final = tuple(download(gateway, key, file_id) for gateway in rig.gateways(2))
assert _statuses(downloads) == [200, 200], [response.text for response in downloads]
assert_download_limited(
timed(partial(download, rig.candidate, key, file_id)), minute_ends, file_id, 2, holder, source
)
requests: Final = rig.provider.drain()
assert seen(requests, over) == ()
assert len(uploads_seen(requests, mark)) == 2
assert len(downloads_seen(requests, file_id)) == 2
def test_limited_key_still_reaches_chat_completions_and_file_metadata(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 1, DOWNLOADS: 1})
mark: Final = marker()
file_id: Final = f"file-{mark}"
content: Final = batch_file(mark, 1)
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
assert _accepted(upload(rig.candidate, key, content)) == file_id
assert_upload_limited(timed(partial(upload, rig.sibling, key, content)), day_ends, 1, THIS_KEY, IN_KEY)
assert download(rig.candidate, key, file_id).status_code == 200
assert_download_limited(
timed(partial(download, rig.sibling, key, file_id)), minute_ends, file_id, 1, THIS_KEY, IN_KEY
)
for gateway in rig.gateways(4):
completion: Final = gateway.chat(ROUTED_MODEL, key=key, text=mark)
assert completion["choices"][0]["message"]["content"] == COMPLETION_TEXT, completion
described: Final = gateway.request("GET", f"/v1/files/{file_id}", key=key)
assert described.status_code == 200, described.text
assert described.json()["id"] == file_id, described.text
def _lowered_probe(rig: Rig, key: str) -> tuple[httpx.Response, httpx.Response, httpx.Response]:
file_id: Final = f"file-{marker()}"
return (
upload(rig.sibling, key, batch_file(marker(), 2)),
download(rig.sibling, key, file_id),
download(rig.sibling, key, file_id),
)
@pytest.mark.parametrize("level", LEVELS)
def test_a_lowered_limit_reaches_the_other_proxy_process(rig: Rig, level: str) -> None:
with rig.candidate.scenario() as scenario:
team: Final = scenario.team(metadata={RECORDS: 3, DOWNLOADS: 5}) if level == "team" else None
key: Final = _key(
scenario, team, **({UPLOADS: ROOMY} if level == "team" else {UPLOADS: ROOMY, RECORDS: 3, DOWNLOADS: 5})
)
source: Final = IN_KEY if level == "key" else IN_TEAM
assert _statuses(_lowered_probe(rig, key)) == [200, 200, 200]
if team is None:
rig.candidate.post("/key/update", {"key": key, "metadata": {UPLOADS: ROOMY, RECORDS: 1, DOWNLOADS: 1}})
else:
rig.candidate.post("/team/update", {"team_id": team, "metadata": {RECORDS: 1, DOWNLOADS: 1}})
too_long, _, _ = eventually(
partial(_lowered_probe, rig, key), lambda probe: _statuses(probe) == [413, 200, 429], seconds=30
)
assert_too_many_records(too_long, 1, source)
@asynccontextmanager
async def _async_clients(rig: Rig, count: int) -> AsyncIterator[tuple[httpx.AsyncClient, ...]]:
async with (
httpx.AsyncClient(base_url=rig.candidate.client.base_url, timeout=60, trust_env=False) as first,
httpx.AsyncClient(base_url=rig.sibling.client.base_url, timeout=60, trust_env=False) as second,
):
yield tuple((first, second)[index % 2] for index in range(count))
def _raise_then_lower(rig: Rig, key: str) -> None:
rig.candidate.post("/key/update", {"key": key, "metadata": {DOWNLOADS: 8}})
rig.candidate.post("/key/update", {"key": key, "metadata": {DOWNLOADS: 3}})
async def test_a_limit_changed_during_a_download_burst_only_answers_200_or_429(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: 5})
file_id: Final = f"file-{marker()}"
headers: Final = {"Authorization": f"Bearer {key}"}
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
async with _async_clients(rig, 30) as clients:
*responses, _ = await asyncio.gather(
*(client.get(f"/v1/files/{file_id}/content", headers=headers) for client in clients),
asyncio.to_thread(_raise_then_lower, rig, key),
)
assert time.time() < minute_ends
statuses: Final = [response.status_code for response in responses]
assert set(statuses) <= {200, 429}, [response.text for response in responses]
assert 3 <= statuses.count(200) <= 8, statuses
assert len(downloads_seen(rig.provider.drain(), file_id)) == statuses.count(200)
async def test_concurrent_requests_across_processes_never_exceed_a_limit(rig: Rig) -> None:
with rig.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: 5, DOWNLOADS: 4})
mark: Final = marker()
file_id: Final = f"file-{mark}"
headers: Final = {"Authorization": f"Bearer {key}"}
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
async with _async_clients(rig, 24) as clients:
uploads: Final = await asyncio.gather(
*(
client.post(
"/v1/files",
data={"purpose": "batch"},
files={"file": ("batch.jsonl", batch_file(mark, 1), "application/jsonl")},
headers=headers,
)
for client in clients
)
)
downloads: Final = await asyncio.gather(
*(client.get(f"/v1/files/{file_id}/content", headers=headers) for client in clients[:20])
)
assert time.time() < min(day_ends, minute_ends)
assert sorted(response.status_code for response in uploads) == [200] * 5 + [429] * 19
assert sorted(response.status_code for response in downloads) == [200] * 4 + [429] * 16
requests: Final = rig.provider.drain()
assert len(uploads_seen(requests, mark)) == 5
assert len(downloads_seen(requests, file_id)) == 4
assert _counts(rig, UPLOADS, hashed(key)) == [5]
assert _counts(rig, DOWNLOADS, f"{hashed(key)}:{file_id}:") == [4]

View file

@ -0,0 +1,411 @@
import asyncio
import re
import signal
import time
from collections.abc import Callable, Coroutine, Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from functools import partial
from pathlib import Path
from queue import SimpleQueue
from threading import Event
from typing import Final
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, string_value
from integration._support.database import scratch_database
from integration._support.process import owned_proxy, owned_proxy_process
from integration._support.redis_process import OwnedRedis, owned_redis
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.management._batch_file_caps import (
DAY_SECONDS,
DOWNLOADS,
IN_GENERAL_SETTINGS,
IN_KEY,
MINUTE_SECONDS,
THIS_KEY,
UPLOADS,
Timed,
assert_download_limited,
assert_upload_limited,
batch_file,
caps_config,
config_entry,
download,
downloads_seen,
file_content,
marker,
provider,
provider_environment,
timed,
upload,
uploads_seen,
window_end,
)
from pydantic import JsonValue
pytestmark = pytest.mark.timeout(600)
WORKERS: Final = 2
UPLOAD_CAP: Final = 5
DOWNLOAD_CAP: Final = 3
HELD_UPLOADS: Final = 20
RESTART_CAP: Final = 4
STORED_CAP: Final = 2
BURST: Final = 20
DAY_ROOM_SECONDS: Final = 180
MINUTE_ROOM_SECONDS: Final = 30
BURST_ROOM_SECONDS: Final = 30
RELOAD_SECONDS: Final = "3"
BREAKER_RECOVERY_SECONDS: Final = "2"
RECOVERY_SECONDS: Final = 60
STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
@dataclass(frozen=True, slots=True)
class Chaos:
candidate: Gateway
sibling: Gateway
provider: Wire
redis: OwnedRedis
config: Path
environment: Mapping[str, str]
directory: Path
@pytest.fixture(scope="module")
def chaos(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Chaos]:
directory: Final = tmp_path_factory.mktemp("batch_file_caps_chaos")
with (
gateway_from_environment() as gateway,
wire_server(provider) as files,
owned_redis(directory) as redis,
):
config: Final = caps_config(directory, files.url, {})
environment: Final = {
**provider_environment(files.url),
"REDIS_HOST": redis.host,
"REDIS_PORT": str(redis.port),
"REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": BREAKER_RECOVERY_SECONDS,
}
with (
owned_proxy(gateway, directory, environment, config=config, workers=WORKERS) as candidate,
owned_proxy(gateway, directory, environment, config=config) as sibling,
):
yield Chaos(candidate, sibling, files, redis, config, environment, directory)
def _key(scenario: Scenario, **limits: JsonValue) -> str:
return scenario.key(metadata=limits)
def _statuses(responses: tuple[httpx.Response, ...]) -> list[int]:
return sorted(response.status_code for response in responses)
def _accepted(response: httpx.Response) -> str:
assert response.status_code == 200, response.text
return string_value(response.json()["id"])
async def _upload_burst(
base_url: str, key: str, mark: str, count: int, *, tolerate_disconnects: bool = False
) -> tuple[httpx.Response, ...]:
headers: Final = {"Authorization": f"Bearer {key}"}
async with httpx.AsyncClient(base_url=base_url, timeout=90, trust_env=False) as client:
results: Final = await asyncio.gather(
*(
client.post(
"/v1/files",
data={"purpose": "batch"},
files={"file": ("batch.jsonl", batch_file(mark, 1), "application/jsonl")},
headers=headers,
)
for _ in range(count)
),
return_exceptions=tolerate_disconnects,
)
for result in results:
assert isinstance(result, httpx.Response | httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, httpx.Response))
async def _download_burst(base_url: str, key: str, file_id: str, count: int) -> tuple[httpx.Response, ...]:
headers: Final = {"Authorization": f"Bearer {key}"}
async with httpx.AsyncClient(base_url=base_url, timeout=90, trust_env=False) as client:
return await asyncio.gather(
*(client.get(f"/v1/files/{file_id}/content", headers=headers) for _ in range(count))
)
@dataclass(frozen=True, slots=True)
class Burst:
subject: str
window_ends: float
answers: tuple[Timed, ...]
def statuses(self) -> list[int]:
return _statuses(tuple(observed.response for observed in self.answers))
def refused(self) -> tuple[Timed, ...]:
return tuple(observed for observed in self.answers if observed.response.status_code == 429)
def allowed_exactly(self, count: int) -> bool:
in_window: Final = all(observed.after < self.window_ends for observed in self.answers)
return in_window and self.statuses() == [200] * count + [429] * (len(self.answers) - count)
def _timed_burst(burst: Coroutine[object, None, tuple[httpx.Response, ...]]) -> tuple[Timed, ...]:
before: Final = time.time()
responses: Final = asyncio.run(burst)
after: Final = time.time()
return tuple(Timed(response, before, after) for response in responses)
def _uploads_after_the_sibling_used_the_cap(chaos: Chaos, scenario: Scenario) -> Burst:
key: Final = _key(scenario, **{UPLOADS: UPLOAD_CAP})
mark: Final = marker()
day_ends: Final = window_end(DAY_SECONDS, BURST_ROOM_SECONDS)
counted: Final = tuple(upload(chaos.sibling, key, batch_file(mark, 1)) for _ in range(UPLOAD_CAP))
assert _statuses(counted) == [200] * UPLOAD_CAP, [response.text for response in counted]
return Burst(mark, day_ends, _timed_burst(_upload_burst(str(chaos.candidate.client.base_url), key, mark, BURST)))
def _downloads_after_the_sibling_used_the_cap(chaos: Chaos, scenario: Scenario) -> Burst:
key: Final = _key(scenario, **{DOWNLOADS: DOWNLOAD_CAP})
file_id: Final = f"file-{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, BURST_ROOM_SECONDS)
counted: Final = tuple(download(chaos.sibling, key, file_id) for _ in range(DOWNLOAD_CAP))
assert _statuses(counted) == [200] * DOWNLOAD_CAP, [response.text for response in counted]
assert {response.content for response in counted} == {file_content(file_id)}
return Burst(
file_id, minute_ends, _timed_burst(_download_burst(str(chaos.candidate.client.base_url), key, file_id, BURST))
)
async def test_uploads_fall_back_to_per_process_counts_while_redis_is_down(chaos: Chaos) -> None:
with chaos.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: UPLOAD_CAP})
mark: Final = marker()
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
before: Final = tuple(upload(chaos.candidate, key, batch_file(mark, 1)) for _ in range(UPLOAD_CAP - 2))
assert _statuses(before) == [200] * (UPLOAD_CAP - 2), [response.text for response in before]
await asyncio.to_thread(chaos.redis.stop)
try:
during: Final = await _upload_burst(str(chaos.candidate.client.base_url), key, mark, BURST)
finally:
await asyncio.to_thread(chaos.redis.start)
assert time.time() < day_ends
statuses: Final = _statuses(during)
assert set(statuses) <= {200, 429}, [response.text for response in during]
accepted: Final = len(before) + statuses.count(200)
assert UPLOAD_CAP <= accepted <= UPLOAD_CAP * WORKERS, statuses
assert len(uploads_seen(chaos.provider.drain(), mark)) == accepted
shared: Final = await asyncio.to_thread(
eventually,
partial(_uploads_after_the_sibling_used_the_cap, chaos, scenario),
lambda burst: burst.allowed_exactly(0),
RECOVERY_SECONDS,
)
for refusal in shared.refused():
assert_upload_limited(refusal, shared.window_ends, UPLOAD_CAP, THIS_KEY, IN_KEY)
assert len(uploads_seen(chaos.provider.drain(), shared.subject)) == UPLOAD_CAP
async def test_downloads_keep_answering_while_redis_hangs(chaos: Chaos) -> None:
with chaos.candidate.scenario() as scenario:
key: Final = _key(scenario, **{DOWNLOADS: DOWNLOAD_CAP})
file_id: Final = f"file-{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, MINUTE_ROOM_SECONDS)
chaos.redis.signal(signal.SIGSTOP)
try:
burst: Final = asyncio.create_task(
_download_burst(str(chaos.candidate.client.base_url), key, file_id, BURST)
)
liveliness: Final = await asyncio.to_thread(
timed, partial(chaos.candidate.request, "GET", "/health/liveliness")
)
during: Final = await burst
finally:
chaos.redis.signal(signal.SIGCONT)
assert time.time() < minute_ends
assert liveliness.response.status_code == 200, liveliness.response.text
assert liveliness.after - liveliness.before < 5, liveliness
statuses: Final = _statuses(during)
assert set(statuses) <= {200, 429}, [response.text for response in during]
assert DOWNLOAD_CAP <= statuses.count(200) <= DOWNLOAD_CAP * WORKERS, statuses
assert len(downloads_seen(chaos.provider.drain(), file_id)) == statuses.count(200)
shared: Final = await asyncio.to_thread(
eventually,
partial(_downloads_after_the_sibling_used_the_cap, chaos, scenario),
lambda burst: burst.allowed_exactly(0),
RECOVERY_SECONDS,
)
for refusal in shared.refused():
assert_download_limited(refusal, shared.window_ends, shared.subject, DOWNLOAD_CAP, THIS_KEY, IN_KEY)
assert len(downloads_seen(chaos.provider.drain(), shared.subject)) == DOWNLOAD_CAP
def _held_provider(release: Event, held: SimpleQueue[str]) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if (request.method, request.target) == ("POST", "/v1/files"):
held.put(request.target)
assert release.wait(timeout=120), "The held uploads were never released"
return provider(request)
return respond
@contextmanager
def _released_on_exit(release: Event) -> Iterator[None]:
try:
yield
finally:
release.set()
def _open_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
def _worker_pids(log: Path) -> tuple[int, ...]:
return tuple(int(pid) for pid in STARTED_WORKER.findall(log.read_text()))
async def test_a_killed_worker_keeps_its_upload_slots_used_and_the_sibling_serving(
chaos: Chaos, tmp_path: Path
) -> None:
release: Final = Event()
held: Final[SimpleQueue[str]] = SimpleQueue()
with wire_server(_held_provider(release, held)) as files:
config: Final = caps_config(tmp_path, files.url, {})
environment: Final = {**chaos.environment, **provider_environment(files.url)}
with (
owned_proxy_process(chaos.candidate, tmp_path, environment, config=config, workers=WORKERS) as owned,
owned.gateway.scenario() as scenario,
_released_on_exit(release),
):
candidate: Final = owned.gateway
key: Final = _key(scenario, **{UPLOADS: HELD_UPLOADS})
mark: Final = marker()
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
workers: Final = eventually(partial(_worker_pids, owned.log), lambda pids: len(pids) == WORKERS, seconds=30)
burst: Final = asyncio.create_task(
_upload_burst(str(candidate.client.base_url), key, mark, HELD_UPLOADS, tolerate_disconnects=True)
)
await asyncio.to_thread(eventually, held.qsize, lambda size: size == HELD_UPLOADS, 90)
held_by: Final = {pid: _open_upstream_connections(pid, files.url) for pid in workers}
assert sum(held_by.values()) == HELD_UPLOADS, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
served: Final = await burst
assert _statuses(served) == [200] * held_by[survivor_pid], (held_by, [response.text for response in served])
assert_upload_limited(
timed(partial(upload, candidate, key, batch_file(marker(), 1))),
day_ends,
HELD_UPLOADS,
THIS_KEY,
IN_KEY,
)
assert len(uploads_seen(files.drain(), mark)) == HELD_UPLOADS
assert candidate.request("GET", "/health/liveliness").status_code == 200
def test_the_daily_count_survives_a_proxy_restart(chaos: Chaos) -> None:
with chaos.candidate.scenario() as scenario:
key: Final = _key(scenario, **{UPLOADS: RESTART_CAP})
mark: Final = marker()
with owned_proxy(chaos.candidate, chaos.directory, chaos.environment, config=chaos.config) as first:
day_ends: Final = window_end(DAY_SECONDS, DAY_ROOM_SECONDS)
before: Final = tuple(upload(first, key, batch_file(mark, 1)) for _ in range(RESTART_CAP - 1))
assert _statuses(before) == [200] * (RESTART_CAP - 1), [response.text for response in before]
with owned_proxy(chaos.candidate, chaos.directory, chaos.environment, config=chaos.config) as second:
assert upload(second, key, batch_file(mark, 1)).status_code == 200
assert_upload_limited(
timed(partial(upload, second, key, batch_file(mark, 1))), day_ends, RESTART_CAP, THIS_KEY, IN_KEY
)
assert len(uploads_seen(chaos.provider.drain(), mark)) == RESTART_CAP
def _plain_key(gateway: Gateway) -> str:
return string_value(gateway.post("/key/generate", {})["key"])
def _download_statuses(gateway: Gateway, key: str) -> list[int]:
file_id: Final = f"file-{marker()}"
responses: Final = asyncio.run(_download_burst(str(gateway.client.base_url), key, file_id, BURST))
return _statuses(responses)
def _downloads_in_one_minute(gateway: Gateway, key: str) -> Burst:
file_id: Final = f"file-{marker()}"
minute_ends: Final = window_end(MINUTE_SECONDS, BURST_ROOM_SECONDS)
return Burst(file_id, minute_ends, _timed_burst(_download_burst(str(gateway.client.base_url), key, file_id, BURST)))
def _assert_stored_cap(burst: Burst) -> None:
assert burst.allowed_exactly(STORED_CAP), burst.statuses()
for refusal in burst.refused():
assert_download_limited(refusal, burst.window_ends, burst.subject, STORED_CAP, THIS_KEY, IN_GENERAL_SETTINGS)
def _stored_cap_on_every_worker(gateway: Gateway, key: str) -> Burst:
return eventually(
partial(_downloads_in_one_minute, gateway, key), lambda burst: burst.allowed_exactly(STORED_CAP), seconds=30
)
def test_a_download_cap_stored_through_the_config_api_reaches_every_worker_and_survives_a_restart(
chaos: Chaos, tmp_path: Path
) -> None:
with scratch_database() as database_url:
environment: Final = {
**chaos.environment,
"DATABASE_URL": database_url,
"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": RELOAD_SECONDS,
}
boot: Final = partial(
owned_proxy,
chaos.candidate,
tmp_path,
environment,
config=chaos.config,
remove_environment=("DATABASE_URL_READ_REPLICA",),
)
with boot(workers=WORKERS) as candidate:
first: Final = _plain_key(candidate)
second: Final = _plain_key(candidate)
assert _download_statuses(candidate, first) == [200] * BURST
assert config_entry(candidate, DOWNLOADS)["field_value"] is None
stored: Final = candidate.request(
"POST",
"/config/field/update",
{"field_name": DOWNLOADS, "field_value": STORED_CAP, "config_type": "general_settings"},
)
assert stored.status_code == 200, stored.text
entry: Final = config_entry(candidate, DOWNLOADS)
assert (entry["field_value"], entry["stored_in_db"]) == (STORED_CAP, True), entry
_assert_stored_cap(_stored_cap_on_every_worker(candidate, first))
_assert_stored_cap(_stored_cap_on_every_worker(candidate, second))
with boot() as restarted:
_assert_stored_cap(_downloads_in_one_minute(restarted, first))
assert config_entry(restarted, DOWNLOADS)["field_value"] == STORED_CAP
removed: Final = restarted.request(
"POST", "/config/field/delete", {"field_name": DOWNLOADS, "config_type": "general_settings"}
)
assert removed.status_code == 200, removed.text
eventually(
partial(_download_statuses, restarted, second), lambda statuses: statuses == [200] * BURST, seconds=30
)
assert config_entry(restarted, DOWNLOADS)["field_value"] is None

View file

@ -0,0 +1,390 @@
import uuid
from collections.abc import Mapping
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value
from integration._support.database import read_rows
from integration.management._batch_file_caps import DOWNLOADS, RECORDS, UPLOADS, config_entry, hashed
from pydantic import JsonValue
TOKENS: Final = "batch_enqueued_token_limit"
LIMITS: Final = (TOKENS, RECORDS, UPLOADS, DOWNLOADS)
FILE_LIMITS: Final = (RECORDS, UPLOADS, DOWNLOADS)
EVERY_LIMIT: Final[Mapping[str, JsonValue]] = {TOKENS: 9000, RECORDS: 7, UPLOADS: 6, DOWNLOADS: 5}
BULK_ROUTE: Final = "/management/v1/users/bulk"
def _refusal(setting: str, entity: str) -> str:
return (
f"Only proxy admins can set {setting} on a {entity}. "
"It limits what the holder can do with batches, so the holder cannot raise it."
)
def _assert_refused(response: httpx.Response, setting: str, entity: str) -> None:
assert response.status_code == 403, response.text
error: Final = object_value(response.json()["error"])
assert error["message"] == str({"error": _refusal(setting, entity)}), response.text
assert error["code"] == "403", response.text
def _key_metadata(token: str) -> JsonValue:
rows: Final = read_rows('SELECT metadata FROM "LiteLLM_VerificationToken" WHERE token = %s', (hashed(token),))
assert len(rows) == 1, rows
return rows[0]["metadata"]
def _team_metadata(team: str) -> JsonValue:
rows: Final = read_rows('SELECT metadata FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,))
assert len(rows) == 1, rows
return rows[0]["metadata"]
def _user_metadata(user: str) -> JsonValue:
rows: Final = read_rows('SELECT metadata FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,))
assert len(rows) == 1, rows
return rows[0]["metadata"]
def _user_key_metadata(user: str) -> list[JsonValue]:
rows: Final = read_rows('SELECT metadata FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,))
return [row["metadata"] for row in rows]
def _team_admin_key(scenario: Scenario, team: str) -> str:
return scenario.key(user_id=scenario.member(team, role="admin"), team_id=team)
def _org_admin_key(scenario: Scenario, organization: str) -> str:
return scenario.key(user_id=scenario.org_member(organization, role="org_admin"))
def _replaceable_key(gateway: Gateway, scenario: Scenario, **fields: JsonValue) -> str:
key: Final = string_value(gateway.post("/key/generate", fields)["key"])
scenario.cleanups.callback(delete_key_if_present, gateway, key)
return key
def _drop_created_key(scenario: Scenario, response: httpx.Response) -> None:
if response.status_code == 200 and response.json().get("key"):
scenario.cleanups.callback(scenario.delete_key, string_value(response.json()["key"]))
def _drop_created_user(scenario: Scenario, response: httpx.Response, user: str) -> None:
if response.status_code == 200:
scenario.cleanups.callback(scenario.delete_user, user)
_drop_created_key(scenario, response)
def _drop_bulk_rows(scenario: Scenario, response: httpx.Response) -> None:
if response.status_code != 200:
return
for row in response.json()["data"]:
if row["success"]:
scenario.cleanups.callback(scenario.delete_user, string_value(row["user_id"]))
if row["key"]:
scenario.cleanups.callback(scenario.delete_key, string_value(row["key"]))
@pytest.mark.parametrize("setting", LIMITS)
def test_team_admin_cannot_put_a_batch_limit_on_a_new_key(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
caller: Final = _team_admin_key(scenario, team)
alias: Final = f"batch-limit-{uuid.uuid4().hex}"
refused: Final = gateway.request(
"POST", "/key/generate", {"team_id": team, "key_alias": alias, "metadata": {setting: 5}}, key=caller
)
_drop_created_key(scenario, refused)
_assert_refused(refused, setting, "key")
assert read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', (alias,)) == []
plain: Final = gateway.request("POST", "/key/generate", {"team_id": team, "key_alias": alias}, key=caller)
_drop_created_key(scenario, plain)
assert plain.status_code == 200, plain.text
assert _key_metadata(string_value(plain.json()["key"])) == {}
@pytest.mark.parametrize("setting", LIMITS)
def test_team_admin_cannot_put_a_batch_limit_on_an_existing_key(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
caller: Final = _team_admin_key(scenario, team)
key: Final = scenario.key(team_id=team, metadata={"owner": "batch-limit-audit"})
refused: Final = gateway.request(
"POST", "/key/update", {"key": key, "metadata": {"owner": "batch-limit-audit", setting: 5}}, key=caller
)
_assert_refused(refused, setting, "key")
assert _key_metadata(key) == {"owner": "batch-limit-audit"}
@pytest.mark.parametrize("setting", LIMITS)
def test_team_admin_cannot_put_a_batch_limit_on_a_regenerated_key(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
caller: Final = _team_admin_key(scenario, team)
key: Final = _replaceable_key(gateway, scenario, team_id=team)
refused: Final = gateway.request("POST", "/key/regenerate", {"key": key, "metadata": {setting: 5}}, key=caller)
_drop_created_key(scenario, refused)
_assert_refused(refused, setting, "key")
assert _key_metadata(key) == {}
@pytest.mark.parametrize("setting", LIMITS)
def test_org_admin_cannot_put_a_batch_limit_on_an_existing_team(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
organization: Final = scenario.organization()
caller: Final = _org_admin_key(scenario, organization)
team: Final = scenario.team(organization_id=organization, metadata={"owner": "batch-limit-audit"})
refused: Final = gateway.request(
"POST",
"/team/update",
{"team_id": team, "metadata": {"owner": "batch-limit-audit", setting: 5}},
key=caller,
)
_assert_refused(refused, setting, "team")
assert _team_metadata(team) == {"owner": "batch-limit-audit"}
@pytest.mark.parametrize("setting", LIMITS)
def test_org_admin_cannot_put_a_batch_limit_on_a_new_team(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
organization: Final = scenario.organization()
caller: Final = _org_admin_key(scenario, organization)
alias: Final = f"batch-limit-{uuid.uuid4().hex}"
refused: Final = gateway.request(
"POST",
"/team/new",
{"team_alias": alias, "organization_id": organization, "metadata": {setting: 5}},
key=caller,
)
if refused.status_code == 200:
scenario.cleanups.callback(scenario.delete_team, string_value(refused.json()["team_id"]))
_assert_refused(refused, setting, "team")
assert read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_alias = %s', (alias,)) == []
plain: Final = gateway.request(
"POST", "/team/new", {"team_alias": alias, "organization_id": organization}, key=caller
)
if plain.status_code == 200:
scenario.cleanups.callback(scenario.delete_team, string_value(plain.json()["team_id"]))
assert plain.status_code == 200, plain.text
assert _team_metadata(string_value(plain.json()["team_id"])) == {}
@pytest.mark.parametrize("setting", LIMITS)
def test_org_admin_cannot_put_a_batch_limit_on_a_new_users_key(gateway: Gateway, setting: str) -> None:
with gateway.scenario() as scenario:
organization: Final = scenario.organization()
caller: Final = _org_admin_key(scenario, organization)
refused_user, keyless_user, plain_user = (f"batch-limit-{uuid.uuid4().hex}" for _ in range(3))
refused: Final = gateway.request(
"POST",
"/user/new",
{"user_id": refused_user, "organization_id": organization, "metadata": {setting: 5}},
key=caller,
)
_drop_created_user(scenario, refused, refused_user)
_assert_refused(refused, setting, "key")
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (refused_user,)) == []
assert _user_key_metadata(refused_user) == []
keyless: Final = gateway.request(
"POST",
"/user/new",
{
"user_id": keyless_user,
"organization_id": organization,
"metadata": {setting: 5},
"auto_create_key": False,
},
key=caller,
)
_drop_created_user(scenario, keyless, keyless_user)
assert keyless.status_code == 200, keyless.text
assert _user_metadata(keyless_user) == {setting: 5}
assert _user_key_metadata(keyless_user) == []
plain: Final = gateway.request(
"POST", "/user/new", {"user_id": plain_user, "organization_id": organization}, key=caller
)
_drop_created_user(scenario, plain, plain_user)
assert plain.status_code == 200, plain.text
assert _user_key_metadata(plain_user) == [{}]
@pytest.mark.parametrize("setting", LIMITS)
def test_bulk_user_creation_refuses_only_the_row_that_puts_a_batch_limit_on_a_key(
gateway: Gateway, setting: str
) -> None:
with gateway.scenario() as scenario:
caller: Final = scenario.key(user_id=scenario.user(user_role="internal_user"), allowed_routes=[BULK_ROUTE])
refused_user, keyless_user, plain_user = (f"batch-limit-{uuid.uuid4().hex}" for _ in range(3))
response: Final = gateway.request(
"POST",
BULK_ROUTE,
{
"users": [
{"user_id": refused_user, "metadata": {setting: 5}, "auto_create_key": True},
{"user_id": keyless_user, "metadata": {setting: 5}},
{"user_id": plain_user, "auto_create_key": True},
]
},
key=caller,
)
_drop_bulk_rows(scenario, response)
assert response.status_code == 200, response.text
refused, keyless, plain = response.json()["data"]
assert refused == {
"user_id": refused_user,
"user_email": None,
"success": False,
"teams": None,
"key": None,
"error": _refusal(setting, "key"),
}, response.text
assert (keyless["success"], keyless["key"], keyless["error"]) == (True, None, None), response.text
assert (plain["success"], plain["error"]) == (True, None), response.text
assert response.json()["meta"] == {"total_requested": 3, "created": 2, "failed": 1}, response.text
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (refused_user,)) == []
assert _user_key_metadata(refused_user) == []
assert _user_metadata(keyless_user) == {setting: 5}
assert _user_key_metadata(keyless_user) == []
assert _user_key_metadata(plain_user) == [{}]
assert _key_metadata(string_value(plain["key"])) == {}
def test_proxy_admin_sets_every_batch_limit_on_every_route(gateway: Gateway) -> None:
limits: Final = dict(EVERY_LIMIT)
with gateway.scenario() as scenario:
generated: Final = scenario.key(metadata=limits)
assert _key_metadata(generated) == limits
updated: Final = scenario.key()
gateway.post("/key/update", {"key": updated, "metadata": limits})
assert _key_metadata(updated) == limits
replaced: Final = _replaceable_key(gateway, scenario)
regenerated: Final = string_value(gateway.post("/key/regenerate", {"key": replaced, "metadata": limits})["key"])
scenario.cleanups.callback(scenario.delete_key, regenerated)
assert _key_metadata(regenerated) == limits
created_team: Final = scenario.team(metadata=limits)
assert _team_metadata(created_team) == limits
updated_team: Final = scenario.team()
gateway.post("/team/update", {"team_id": updated_team, "metadata": limits})
assert _team_metadata(updated_team) == limits
new_user, bulk_user = (f"batch-limit-{uuid.uuid4().hex}" for _ in range(2))
created_user: Final = gateway.request("POST", "/user/new", {"user_id": new_user, "metadata": limits})
_drop_created_user(scenario, created_user, new_user)
assert created_user.status_code == 200, created_user.text
assert _user_key_metadata(new_user) == [limits]
bulk: Final = gateway.request(
"POST", BULK_ROUTE, {"users": [{"user_id": bulk_user, "metadata": limits, "auto_create_key": True}]}
)
_drop_bulk_rows(scenario, bulk)
assert bulk.status_code == 200, bulk.text
assert bulk.json()["meta"] == {"total_requested": 1, "created": 1, "failed": 0}, bulk.text
assert _user_key_metadata(bulk_user) == [limits]
RESENDS: Final = (
pytest.param({}, id="dropped"),
pytest.param({UPLOADS: 4}, id="lowered"),
pytest.param({UPLOADS: 6}, id="raised"),
pytest.param({UPLOADS: "5"}, id="resent-as-a-string"),
pytest.param({UPLOADS: None}, id="resent-as-null"),
)
@pytest.mark.parametrize("metadata", RESENDS)
def test_team_admin_may_resend_a_stored_key_limit_but_not_change_it(
gateway: Gateway, metadata: Mapping[str, JsonValue]
) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
caller: Final = _team_admin_key(scenario, team)
key: Final = scenario.key(team_id=team, metadata={UPLOADS: 5})
alias: Final = f"batch-limit-{uuid.uuid4().hex}"
resent: Final = gateway.request(
"POST", "/key/update", {"key": key, "key_alias": alias, "metadata": {UPLOADS: 5}}, key=caller
)
assert resent.status_code == 200, resent.text
untouched: Final = gateway.request("POST", "/key/update", {"key": key, "key_alias": f"{alias}-2"}, key=caller)
assert untouched.status_code == 200, untouched.text
changed: Final = gateway.request("POST", "/key/update", {"key": key, "metadata": dict(metadata)}, key=caller)
_assert_refused(changed, UPLOADS, "key")
assert _key_metadata(key) == {UPLOADS: 5}
assert read_rows('SELECT key_alias FROM "LiteLLM_VerificationToken" WHERE token = %s', (hashed(key),)) == [
{"key_alias": f"{alias}-2"}
]
@pytest.mark.parametrize("metadata", RESENDS)
def test_org_admin_may_resend_a_stored_team_limit_but_not_change_it(
gateway: Gateway, metadata: Mapping[str, JsonValue]
) -> None:
with gateway.scenario() as scenario:
organization: Final = scenario.organization()
caller: Final = _org_admin_key(scenario, organization)
team: Final = scenario.team(organization_id=organization, metadata={UPLOADS: 5})
resent: Final = gateway.request(
"POST", "/team/update", {"team_id": team, "tpm_limit": 5000, "metadata": {UPLOADS: 5}}, key=caller
)
assert resent.status_code == 200, resent.text
untouched: Final = gateway.request("POST", "/team/update", {"team_id": team, "tpm_limit": 6000}, key=caller)
assert untouched.status_code == 200, untouched.text
changed: Final = gateway.request(
"POST", "/team/update", {"team_id": team, "metadata": dict(metadata)}, key=caller
)
_assert_refused(changed, UPLOADS, "team")
assert read_rows('SELECT metadata, tpm_limit FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)) == [
{"metadata": {UPLOADS: 5}, "tpm_limit": 6000}
]
@pytest.mark.parametrize("setting", FILE_LIMITS)
def test_config_list_shows_each_file_limit_as_an_unset_integer(gateway: Gateway, setting: str) -> None:
entry: Final = config_entry(gateway, setting)
assert (entry["field_type"], entry["field_value"], entry["stored_in_db"], entry["editable"]) == (
"Integer",
None,
None,
True,
), entry
@pytest.mark.parametrize("setting", FILE_LIMITS)
def test_config_update_refuses_a_zero_file_limit(gateway: Gateway, setting: str) -> None:
response: Final = gateway.request("POST", "/config/update", {"general_settings": {setting: 0}})
assert response.status_code == 422, response.text
assert response.json() == {
"detail": [
{
"type": "greater_than",
"loc": ["body", "general_settings", setting],
"msg": "Input should be greater than 0",
}
]
}, response.text
assert config_entry(gateway, setting)["field_value"] is None
@pytest.mark.parametrize("setting", FILE_LIMITS)
@pytest.mark.parametrize(
("value", "kind"),
[
pytest.param(0, "int", id="zero"),
pytest.param(-1, "int", id="negative"),
pytest.param("abc", "str", id="word"),
pytest.param(1.5, "float", id="fraction"),
pytest.param([5], "list", id="list"),
pytest.param("", "str", id="empty-string"),
],
)
def test_config_field_update_refuses_a_malformed_file_limit(
gateway: Gateway, value: JsonValue, kind: str, setting: str
) -> None:
response: Final = gateway.request(
"POST",
"/config/field/update",
{"field_name": setting, "field_value": value, "config_type": "general_settings"},
)
assert response.status_code == 400, response.text
assert response.json() == {"detail": {"error": f"Invalid type of field value=<class '{kind}'> passed in."}}
assert config_entry(gateway, setting)["field_value"] is None

View file

@ -0,0 +1,85 @@
import uuid
from collections.abc import Callable, Mapping
from pathlib import Path
from typing import Final
import yaml
from integration._support.client import Gateway, eventually
from integration._support.otlp_sink import Span, SpanSinks, recorded_spans, spans_for_trace
from integration._support.process import owned_proxy
from integration._support.wire import wire_server
from integration.management._batch_file_caps import (
UPLOADS,
batch_file,
caps_config,
marker,
provider,
provider_environment,
uploads_seen,
)
from pydantic import JsonValue
COUNTER_SPAN: Final = "redis.incr rate_limits"
SERVER: Final = 2
CAP: Final = 3
def _config(directory: Path, otel: Path, provider_url: str) -> Path:
tracing: Final = yaml.safe_load(otel.read_text())
caps: Final = yaml.safe_load(caps_config(directory, provider_url, {UPLOADS: CAP}).read_text())
path: Final = directory / "otel-file-usage-caps.yaml"
path.write_text(
yaml.safe_dump(
{
**tracing,
"model_list": caps["model_list"],
"files_settings": caps["files_settings"],
"general_settings": {**tracing["general_settings"], UPLOADS: CAP},
}
)
)
return path
def _names(spans: tuple[Span, ...]) -> list[str]:
return sorted(span["name"] for span in spans)
def _traced_upload(candidate: Gateway, content: bytes, trace_id: str) -> int:
response: Final = candidate.client.post(
"/v1/files",
data={"purpose": "batch"},
files={"file": ("batch.jsonl", content, "application/jsonl")},
headers={
"Authorization": f"Bearer {candidate.key}",
"traceparent": f"00-{trace_id}-{uuid.uuid4().hex[:16]}-01",
},
)
assert response.status_code == 200, response.text
return response.status_code
def test_a_counted_upload_exports_its_counter_increment_as_a_rate_limits_span(
gateway: Gateway,
tmp_path: Path,
audit_sinks: SpanSinks,
otel_audit_config: Callable[[Path, Mapping[str, JsonValue]], Path],
) -> None:
with wire_server(provider) as files:
config: Final = _config(tmp_path, otel_audit_config(tmp_path, {}), files.url)
environment: Final = {**provider_environment(files.url), "LITELLM_OTEL_V2": "1"}
with owned_proxy(gateway, tmp_path, environment, config=config) as candidate:
since, _ = recorded_spans(audit_sinks.operator)
mark: Final = marker()
trace_id: Final = uuid.uuid4().hex
assert _traced_upload(candidate, batch_file(mark, 1), trace_id) == 200
trace: Final = eventually(
lambda: spans_for_trace(recorded_spans(audit_sinks.operator, since)[1], trace_id),
lambda spans: COUNTER_SPAN in _names(spans),
seconds=40,
)
assert _names(trace).count(COUNTER_SPAN) == 1, _names(trace)
assert sum(1 for span in trace if span["kind"] == SERVER) == 1, _names(trace)
(counter,) = (span for span in trace if span["name"] == COUNTER_SPAN)
assert counter["attributes"].get("litellm.service.target") == "rate_limits", counter["attributes"]
assert len(uploads_seen(files.drain(), mark)) == 1

View file

@ -1237,6 +1237,69 @@ async def test_new_user_admin_can_set_permissions(mocker):
assert result is not None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"limit_key",
["max_batch_file_records", "max_batch_file_uploads_per_day", "max_file_downloads_per_minute"],
)
async def test_new_user_only_proxy_admin_sets_batch_limits_on_the_created_key(mocker, limit_key):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
async def mock_count(*args, **kwargs):
return 5
mock_prisma_client.db.litellm_usertable.count = mock_count
async def mock_check(*_args, **_kwargs):
return None
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email",
mock_check,
)
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id",
mock_check,
)
mock_license_check = mocker.MagicMock()
mock_license_check.is_over_limit.return_value = False
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check)
created_with: list[dict[str, object]] = []
async def stub_helper(**kwargs):
created_with.append(kwargs)
return {"user_id": "alice", "key": "sk-alice", "expires": None}
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn",
stub_helper,
)
org_admin = UserAPIKeyAuth(user_id="org-admin", user_role=LitellmUserRoles.ORG_ADMIN)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
def request(auto_create_key: bool) -> NewUserRequest:
return NewUserRequest(
user_email="alice@example.com",
user_role=LitellmUserRoles.INTERNAL_USER,
metadata={limit_key: 1000},
auto_create_key=auto_create_key,
)
with pytest.raises(ProxyException) as exc_info:
await new_user(data=request(auto_create_key=True), user_api_key_dict=org_admin)
assert str(exc_info.value.code) == "403"
assert f"Only proxy admins can set {limit_key} on a key" in str(exc_info.value.message)
assert created_with == []
await new_user(data=request(auto_create_key=False), user_api_key_dict=org_admin)
await new_user(data=request(auto_create_key=True), user_api_key_dict=admin)
assert [call["metadata"] for call in created_with] == [{limit_key: 1000}, {limit_key: 1000}]
@pytest.mark.asyncio
async def test_update_single_user_non_admin_permissions_rejected(mocker):
"""`_update_single_user_helper` rejects a non-admin when `permissions`

View file

@ -18982,30 +18982,37 @@ async def test_regenerate_key_output_token_estimate_lowered_rejected_for_non_adm
_BATCH_LIMIT = "batch_enqueued_token_limit"
_UNTOUCHED = object()
@pytest.mark.parametrize(
"label, request_body, existing_metadata, allowed",
"limit_key",
[
("set on a key with none stored", {"metadata": {_BATCH_LIMIT: 50000}}, None, False),
("raised above the stored limit", {"metadata": {_BATCH_LIMIT: 200000}}, {_BATCH_LIMIT: 100000}, False),
("cleared by replacing the blob", {"metadata": {}}, {_BATCH_LIMIT: 100000}, False),
("resent unchanged", {"metadata": {_BATCH_LIMIT: 100000}}, {_BATCH_LIMIT: 100000}, True),
("left untouched", {}, {_BATCH_LIMIT: 100000}, True),
"batch_enqueued_token_limit",
"max_batch_file_records",
"max_batch_file_uploads_per_day",
"max_file_downloads_per_minute",
],
)
def test_batch_enqueued_token_limit_admin_gate_matrix(label, request_body, existing_metadata, allowed):
"""A non-admin may only leave a key's stored batch enqueued-token limit as it is.
When set, the limit replaces the standard RPM/TPM checks for batch
submissions, so a key holder writing it would pick their own batch quota.
Resending the stored value is what the edit form produces on every save
and has to stay allowed.
"""
@pytest.mark.parametrize(
"label, sent, stored, allowed",
[
("set on a key with none stored", 50000, None, False),
("raised above the stored limit", 200000, 100000, False),
("cleared by replacing the blob", None, 100000, False),
("resent unchanged", 100000, 100000, True),
("left untouched", _UNTOUCHED, 100000, True),
],
)
def test_batch_limits_admin_gate_matrix(limit_key, label, sent, stored, allowed):
request_body = {} if sent is _UNTOUCHED else {"metadata": {} if sent is None else {limit_key: sent}}
existing_metadata = None if stored is None else {limit_key: stored}
from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
enforce_batch_limits_are_admin_only,
)
def _call(caller):
enforce_batch_enqueued_token_limit_is_admin_only(
enforce_batch_limits_are_admin_only(
data=UpdateKeyRequest(key="sk-1", **request_body),
existing_metadata=existing_metadata,
user_api_key_dict=caller,
@ -19023,7 +19030,7 @@ def test_batch_enqueued_token_limit_admin_gate_matrix(label, request_body, exist
with pytest.raises(HTTPException) as exc:
_call(non_admin)
assert exc.value.status_code == 403
assert "Only proxy admins can set" in str(exc.value.detail)
assert f"Only proxy admins can set {limit_key}" in str(exc.value.detail)
_call(
UserAPIKeyAuth(

View file

@ -390,6 +390,37 @@ async def test_non_admin_cannot_create_admin_users_but_other_rows_proceed():
assert set(prisma.db.litellm_usertable.rows) == {"u2"}
@pytest.mark.asyncio
async def test_only_proxy_admin_sets_batch_limits_on_a_created_key():
limits = {"max_file_downloads_per_minute": 1000}
calls: list[dict[str, object]] = []
async def generate_key(**kwargs: object) -> dict[str, object]:
calls.append(kwargs)
return {"token": f"sk-{kwargs['user_id']}"}
rows = [
{"user_id": "u1", "auto_create_key": True, "metadata": limits},
{"user_id": "u2", "auto_create_key": False, "metadata": limits},
{"user_id": "u3", "auto_create_key": True},
]
prisma = _FakePrisma()
response = await _run(prisma, rows, caller=INTERNAL, generate_key=generate_key)
assert [r.success for r in response.data] == [False, True, True]
assert "Only proxy admins can set max_file_downloads_per_minute on a key" in (response.data[0].error or "")
assert set(prisma.db.litellm_usertable.rows) == {"u2", "u3"}
assert [call["user_id"] for call in calls] == ["u3"]
admin_prisma = _FakePrisma()
admin_response = await _run(admin_prisma, rows[:1], caller=ADMIN, generate_key=generate_key)
assert [r.key for r in admin_response.data] == ["sk-u1"]
assert calls[-1]["user_id"] == "u1"
assert calls[-1]["metadata"] == limits
@pytest.mark.asyncio
async def test_license_is_checked_once_against_the_whole_batch():
prisma = _FakePrisma()

View file

@ -0,0 +1,239 @@
from typing import Final
import pytest
from litellm._internal_context import current_service_target
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.openai_files_endpoints.file_usage_caps import (
FileUsageLimit,
ScopedFileUsageLimit,
batch_file_record_limit,
consume_file_usage,
enforce_batch_file_upload_limit,
enforce_file_download_limit,
resolve_scoped_limits,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
DAY: Final = 86400
MIDDAY: Final = 20_000 * DAY + DAY / 2
def _cache() -> InternalUsageCache:
return InternalUsageCache(dual_cache=DualCache())
def _scoped(scope, scope_id, value, source="key", setting="max_batch_file_uploads_per_day"):
return ScopedFileUsageLimit(
scope=scope, scope_id=scope_id, limit=FileUsageLimit(setting=setting, value=value, source=source)
)
@pytest.mark.parametrize(
"key_metadata, team_metadata, general_settings, expected",
[
({}, {}, {}, None),
({}, {}, {"max_batch_file_records": 50}, FileUsageLimit("max_batch_file_records", 50, "general_settings")),
(
{"max_batch_file_records": 80},
{},
{"max_batch_file_records": 50},
FileUsageLimit("max_batch_file_records", 80, "key"),
),
(
{"max_batch_file_records": 80},
{"max_batch_file_records": 30},
{},
FileUsageLimit("max_batch_file_records", 30, "team"),
),
(
{},
{"max_batch_file_records": 90},
{"max_batch_file_records": 50},
FileUsageLimit("max_batch_file_records", 50, "general_settings"),
),
({"max_batch_file_records": 0}, {"max_batch_file_records": "lots"}, {}, None),
({"max_batch_file_records": "40"}, {}, {}, FileUsageLimit("max_batch_file_records", 40, "key")),
],
)
def test_record_limit_is_the_lowest_applicable_limit(key_metadata, team_metadata, general_settings, expected):
caller: Final = UserAPIKeyAuth(api_key="hashed", team_id="t1", metadata=key_metadata, team_metadata=team_metadata)
assert batch_file_record_limit(caller, general_settings) == expected
@pytest.mark.parametrize(
"caller, general_settings, expected",
[
(UserAPIKeyAuth(api_key="hashed", team_id="t1"), {}, ()),
(
UserAPIKeyAuth(api_key="hashed", team_id="t1"),
{"max_batch_file_uploads_per_day": 5},
(_scoped("key", "hashed", 5, "general_settings"),),
),
(
UserAPIKeyAuth(
api_key="hashed",
team_id="t1",
metadata={"max_batch_file_uploads_per_day": 9},
team_metadata={"max_batch_file_uploads_per_day": 20},
),
{"max_batch_file_uploads_per_day": 5},
(_scoped("key", "hashed", 9, "key"), _scoped("team", "t1", 20, "team")),
),
(
UserAPIKeyAuth(api_key="hashed", team_metadata={"max_batch_file_uploads_per_day": 20}),
{},
(),
),
(
UserAPIKeyAuth(api_key=None, user_id="u1", team_id="t1"),
{"max_batch_file_uploads_per_day": 5},
(_scoped("user", "u1", 5, "general_settings"),),
),
(UserAPIKeyAuth(api_key=None), {"max_batch_file_uploads_per_day": 5}, ()),
],
)
def test_counter_scopes_pair_each_limit_with_its_own_counter(caller, general_settings, expected):
assert resolve_scoped_limits(caller, general_settings, "max_batch_file_uploads_per_day") == expected
async def test_a_scope_admits_exactly_its_limit_per_window_and_resets_on_the_next():
cache: Final = _cache()
limits: Final = (_scoped("key", "hashed", 2),)
outcomes: Final = [await consume_file_usage(cache, limits, DAY, "", MIDDAY + i) for i in range(3)]
next_window: Final = await consume_file_usage(cache, limits, DAY, "", MIDDAY + DAY)
assert outcomes[:2] == [None, None]
assert outcomes[2] is not None
assert outcomes[2].limit == limits[0]
assert outcomes[2].retry_after_seconds == DAY / 2 - 2
assert next_window is None
async def test_a_rejection_on_one_scope_gives_back_the_slot_it_took_on_the_other():
cache: Final = _cache()
key_scope: Final = _scoped("key", "hashed", 3)
team_scope: Final = _scoped("team", "t1", 1, "team")
first: Final = await consume_file_usage(cache, (key_scope, team_scope), DAY, "", MIDDAY)
rejected: Final = [await consume_file_usage(cache, (key_scope, team_scope), DAY, "", MIDDAY) for _ in range(5)]
key_only: Final = [await consume_file_usage(cache, (key_scope,), DAY, "", MIDDAY) for _ in range(3)]
assert first is None
assert {outcome.limit for outcome in rejected if outcome is not None} == {team_scope}
assert len([outcome for outcome in rejected if outcome is not None]) == 5
assert key_only[:2] == [None, None]
assert key_only[2] is not None and key_only[2].limit == key_scope
async def test_a_rejection_does_not_use_up_the_scope_that_rejected_it():
cache: Final = _cache()
limits: Final = (_scoped("key", "hashed", 1),)
await consume_file_usage(cache, limits, DAY, "", MIDDAY)
for _ in range(4):
await consume_file_usage(cache, limits, DAY, "", MIDDAY)
raised_limit: Final = (_scoped("key", "hashed", 2),)
assert await consume_file_usage(cache, raised_limit, DAY, "", MIDDAY) is None
class _ServiceTargetSeen(Exception):
pass
class _TargetReportingCache(InternalUsageCache):
async def async_increment_cache(self, key, value, litellm_parent_otel_span, local_only=False, **kwargs):
raise _ServiceTargetSeen(current_service_target())
async def test_counter_writes_are_declared_as_rate_limit_calls_so_their_redis_spans_are_named():
cache: Final = _TargetReportingCache(dual_cache=DualCache())
with pytest.raises(_ServiceTargetSeen) as seen:
await consume_file_usage(cache, (_scoped("key", "hashed", 1),), DAY, "", MIDDAY)
assert seen.value.args == ("rate_limits",)
async def test_download_counters_are_per_file():
cache: Final = _cache()
limits: Final = (_scoped("key", "hashed", 1, setting="max_file_downloads_per_minute"),)
assert await consume_file_usage(cache, limits, 60, "file-a", MIDDAY) is None
assert await consume_file_usage(cache, limits, 60, "file-a", MIDDAY) is not None
assert await consume_file_usage(cache, limits, 60, "file-b", MIDDAY) is None
async def test_the_proxy_counter_store_keeps_a_counter_while_more_counters_are_live_than_a_default_cache_holds():
cache: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()).file_usage_cache
limits: Final = (_scoped("key", "hashed", 1, setting="max_file_downloads_per_minute"),)
other_files: Final = DEFAULT_MAX_SIZE_IN_MEMORY + 10
first: Final = await consume_file_usage(cache, limits, 60, "file-first", MIDDAY)
others: Final = [
await consume_file_usage(cache, limits, 60, f"file-{index}", MIDDAY) for index in range(other_files)
]
assert first is None
assert others == [None] * other_files
assert await consume_file_usage(cache, limits, 60, "file-first", MIDDAY) is not None
async def test_upload_limit_error_names_the_setting_value_scope_and_reset():
cache: Final = _cache()
caller: Final = UserAPIKeyAuth(api_key="hashed", team_id="t1", team_metadata={"max_batch_file_uploads_per_day": 1})
await enforce_batch_file_upload_limit(cache, caller, {}, clock=lambda: MIDDAY)
with pytest.raises(ProxyException) as exc:
await enforce_batch_file_upload_limit(cache, caller, {}, clock=lambda: MIDDAY + 100)
assert exc.value.code == "429"
assert exc.value.type == "rate_limit_error"
assert exc.value.headers == {"retry-after": str(DAY // 2 - 100)}
assert "max_batch_file_uploads_per_day is 1 for team t1" in exc.value.message
assert "this team's metadata" in exc.value.message
async def test_download_limit_error_names_the_file_and_general_settings_default():
cache: Final = _cache()
caller: Final = UserAPIKeyAuth(api_key="hashed")
general_settings: Final = {"max_file_downloads_per_minute": 2}
for _ in range(2):
await enforce_file_download_limit(cache, caller, general_settings, "file-a", clock=lambda: MIDDAY + 15)
with pytest.raises(ProxyException) as exc:
await enforce_file_download_limit(cache, caller, general_settings, "file-a", clock=lambda: MIDDAY + 15)
assert exc.value.code == "429"
assert exc.value.headers == {"retry-after": "45"}
assert "file-a" in exc.value.message
assert "max_file_downloads_per_minute is 2 for this key (set in general_settings)" in exc.value.message
async def test_no_configured_limit_never_rejects():
cache: Final = _cache()
caller: Final = UserAPIKeyAuth(api_key="hashed", team_id="t1")
for _ in range(50):
await enforce_batch_file_upload_limit(cache, caller, {}, clock=lambda: MIDDAY)
await enforce_file_download_limit(cache, caller, {}, "file-a", clock=lambda: MIDDAY)
async def test_a_caller_without_a_key_is_counted_per_user():
cache: Final = _cache()
general_settings: Final = {"max_file_downloads_per_minute": 1}
jwt_caller: Final = UserAPIKeyAuth(api_key=None, user_id="u1")
await enforce_file_download_limit(cache, jwt_caller, general_settings, "file-a", clock=lambda: MIDDAY)
await enforce_file_download_limit(
cache, UserAPIKeyAuth(api_key=None, user_id="u2"), general_settings, "file-a", clock=lambda: MIDDAY
)
with pytest.raises(ProxyException) as exc:
await enforce_file_download_limit(cache, jwt_caller, general_settings, "file-a", clock=lambda: MIDDAY)
assert "max_file_downloads_per_minute is 1 for user u1 (set in general_settings)" in exc.value.message

View file

@ -11,10 +11,12 @@ from litellm.proxy.openai_files_endpoints.batch_file_validation import (
BatchFileLineNotObject,
BatchFileMissingLineKey,
BatchFileTooLarge,
BatchFileTooManyRecords,
BatchFileWrongExtension,
check_batch_file_upload,
raise_batch_file_validation_failure,
)
from litellm.proxy.openai_files_endpoints.file_usage_caps import FileUsageLimit
VALID_LINE = (
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions",'
@ -44,9 +46,7 @@ def test_wrong_extension_rejected(filename):
def test_size_over_cap_rejected_for_bytes():
content = b"x" * (2 * 1024 * 1024)
assert check_batch_file_upload("batch.jsonl", content, 1) == BatchFileTooLarge(
size_bytes=len(content), limit_mb=1
)
assert check_batch_file_upload("batch.jsonl", content, 1) == BatchFileTooLarge(size_bytes=len(content), limit_mb=1)
def test_size_over_cap_rejected_for_binaryio():
@ -150,6 +150,12 @@ def test_scan_stops_at_first_failure():
"file",
("210.0 MB", "max_batch_file_size_mb", "10 MB", "not forwarded"),
),
(
BatchFileTooManyRecords(limit=FileUsageLimit("max_batch_file_records", 1000, "team")),
"413",
"file",
("more than 1000 records", "max_batch_file_records of 1000", "this team's metadata", "not forwarded"),
),
(
BatchFileWrongExtension(filename="batch.csv"),
"400",
@ -208,3 +214,27 @@ def test_passthrough_missing_key_message_says_what_a_passthrough_upload_takes():
assert "passthrough upload takes native Vertex batch rows" in exc_info.value.message
assert "with a request key." in exc_info.value.message
assert "custom_id" not in exc_info.value.message
RECORD_LIMIT_3 = FileUsageLimit("max_batch_file_records", 3, "key")
@pytest.mark.parametrize(
"content, expected",
[
((VALID_LINE + b"\n") * 3, None),
((VALID_LINE + b"\n") * 4, BatchFileTooManyRecords(limit=RECORD_LIMIT_3)),
(b"\n\n" + (VALID_LINE + b"\n\n \n") * 3, None),
(io.BytesIO((VALID_LINE + b"\n") * 4), BatchFileTooManyRecords(limit=RECORD_LIMIT_3)),
(b"broken\n" + (VALID_LINE + b"\n") * 4, BatchFileInvalidJsonLine(line_number=1)),
],
)
def test_record_limit_counts_non_blank_request_lines(content, expected):
assert check_batch_file_upload("batch.jsonl", content, None, max_records=RECORD_LIMIT_3) == expected
def test_record_limit_applies_to_passthrough_rows():
content = (NATIVE_VERTEX_LINE + b"\n") * 4
assert check_batch_file_upload(
"batch.jsonl", content, None, PASSTHROUGH_BATCH_LINE_SHAPE, RECORD_LIMIT_3
) == BatchFileTooManyRecords(limit=RECORD_LIMIT_3)

View file

@ -4224,6 +4224,181 @@ def test_create_file_batch_under_max_batch_file_size_mb_forwards(monkeypatch, ll
assert len(forwarded_calls) == 1
FIXED_NOW: Final = 20_000 * 86400 + 3600 + 15
def _pin_file_usage_clock(monkeypatch) -> None:
from functools import partial
from litellm.proxy.openai_files_endpoints import file_usage_caps
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
monkeypatch.setattr(
fe,
"enforce_batch_file_upload_limit",
partial(file_usage_caps.enforce_batch_file_upload_limit, clock=lambda: FIXED_NOW),
)
monkeypatch.setattr(
fe,
"enforce_file_download_limit",
partial(file_usage_caps.enforce_file_download_limit, clock=lambda: FIXED_NOW),
)
def _upload_batch(content: bytes):
return client.post(
"/v1/files",
files={"file": ("batch.jsonl", content, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
def test_create_file_batch_over_max_batch_file_records_rejected_before_forwarding(monkeypatch, llm_router: Router):
import litellm.proxy.proxy_server as ps
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
monkeypatch.setitem(ps.general_settings, "max_batch_file_records", 2)
try:
over = _upload_batch(VALID_BATCH_LINE * 3)
at_limit = _upload_batch(VALID_BATCH_LINE * 2)
finally:
_teardown_batch_upload_endpoint()
assert over.status_code == 413, over.text
error = over.json()["error"]
assert error["param"] == "file"
assert "max_batch_file_records of 2 set in general_settings" in error["message"]
assert at_limit.status_code == 200, at_limit.text
assert len(forwarded_calls) == 1
def test_create_file_batch_uploads_over_daily_limit_get_429_and_only_valid_files_count(
monkeypatch, llm_router: Router
):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
_pin_file_usage_clock(monkeypatch)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="hashed-caller-key",
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="test-user",
metadata={"max_batch_file_uploads_per_day": 2},
)
try:
invalid = _upload_batch(b"not json\n")
malformed_expiry = client.post(
"/v1/files",
files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")},
data={"purpose": "batch", "expires_after[anchor]": "created_at"},
headers={"Authorization": "Bearer test-key"},
)
allowed = [_upload_batch(VALID_BATCH_LINE) for _ in range(2)]
rejected = _upload_batch(VALID_BATCH_LINE)
finally:
_teardown_batch_upload_endpoint()
assert invalid.status_code == 400, invalid.text
assert malformed_expiry.status_code == 400, malformed_expiry.text
assert [response.status_code for response in allowed] == [200, 200]
assert rejected.status_code == 429, rejected.text
assert rejected.headers["retry-after"] == str(86400 - 3615)
error = rejected.json()["error"]
assert error["type"] == "rate_limit_error"
assert "max_batch_file_uploads_per_day is 2 for this key (set in this key's metadata)" in error["message"]
assert len(forwarded_calls) == 2
def test_get_file_content_over_per_minute_download_limit_gets_429_per_file(monkeypatch, llm_router: Router):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
setup_proxy_logging_object(monkeypatch, llm_router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setitem(ps.general_settings, "max_file_downloads_per_minute", 2)
_pin_file_usage_clock(monkeypatch)
provider_calls: list = []
async def _mock_afile_content(**kwargs):
provider_calls.append(kwargs["file_id"])
async def _stream():
yield b"output"
return FileContentStreamingResult(stream_iterator=_stream(), headers={"content-length": "6"})
monkeypatch.setattr(litellm, "afile_content", _mock_afile_content)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
AsyncMock(return_value=(False, None, None, None)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="hashed-caller-key", user_role=LitellmUserRoles.INTERNAL_USER, user_id="test-user"
)
try:
same_file = [
client.get("/v1/files/file-out/content", headers={"Authorization": "Bearer test-key"}) for _ in range(3)
]
other_file = client.get("/v1/files/file-other/content", headers={"Authorization": "Bearer test-key"})
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert [response.status_code for response in same_file] == [200, 200, 429]
assert same_file[2].headers["retry-after"] == "45"
error = same_file[2].json()["error"]
assert error["type"] == "rate_limit_error"
assert "Download limit reached for file file-out" in error["message"]
assert other_file.status_code == 200, other_file.text
assert provider_calls == ["file-out", "file-out", "file-other"]
def test_download_counters_for_many_files_do_not_evict_rate_limit_state(monkeypatch, llm_router: Router):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
proxy_logging = setup_proxy_logging_object(monkeypatch, llm_router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setitem(ps.general_settings, "max_file_downloads_per_minute", 1000)
_pin_file_usage_clock(monkeypatch)
async def _mock_afile_content(**kwargs):
async def _stream():
yield b"output"
return FileContentStreamingResult(stream_iterator=_stream(), headers={"content-length": "6"})
monkeypatch.setattr(litellm, "afile_content", _mock_afile_content)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
AsyncMock(return_value=(False, None, None, None)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="hashed-caller-key", user_role=LitellmUserRoles.INTERNAL_USER, user_id="test-user"
)
rate_limit_cache: Final = proxy_logging.internal_usage_cache.dual_cache
rate_limit_window_key: Final = "{api_key:hashed-caller-key}:window"
rate_limit_cache.set_cache(key=rate_limit_window_key, value=12345, local_only=True, ttl=60)
distinct_files: Final = rate_limit_cache.in_memory_cache.max_size_in_memory + 1
try:
statuses = [
client.get(f"/v1/files/file-{index}/content", headers={"Authorization": "Bearer test-key"}).status_code
for index in range(distinct_files)
]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert set(statuses) == {200}
assert rate_limit_cache.get_cache(key=rate_limit_window_key, local_only=True) == 12345
def test_create_file_batch_wrong_extension_rejected_before_forwarding(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)

View file

@ -476,3 +476,15 @@ def test_modern_http_upstream_protocol_is_available(request_model):
})
assert parsed.mcp_info["protocol_version"] == "2026-07-28"
assert parsed.transport == "http"
@pytest.mark.parametrize(
"setting", ["max_batch_file_records", "max_batch_file_uploads_per_day", "max_file_downloads_per_minute"]
)
def test_batch_file_caps_accept_only_positive_limits(setting):
from litellm.proxy._types import ConfigGeneralSettings
for invalid in (0, -5):
with pytest.raises(ValidationError):
ConfigGeneralSettings.model_validate({setting: invalid})
assert getattr(ConfigGeneralSettings.model_validate({setting: 3}), setting) == 3

View file

@ -29766,11 +29766,21 @@ export interface components {
* @description require a key for all calls to proxy
*/
master_key?: string | null;
/**
* Max Batch File Records
* @description max records (non-blank lines) per batch input file for /v1/files uploads with purpose=batch, applied per key. A key's metadata can override it and a team's metadata adds a team cap on top, both set by a proxy admin; the lower of the key's value and the team's value wins. Unset means no limit
*/
max_batch_file_records?: number | null;
/**
* Max Batch File Size Mb
* @description max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider
*/
max_batch_file_size_mb?: number | null;
/**
* Max Batch File Uploads Per Day
* @description max /v1/files uploads with purpose=batch per key (per user for JWT callers) per UTC day. A key's metadata can override it and a team's metadata adds a shared team cap, both set by a proxy admin. Unset means no limit
*/
max_batch_file_uploads_per_day?: number | null;
/**
* Max Failed Login Attempts Per Source
* @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded down but at least 1, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10
@ -29783,6 +29793,11 @@ export interface components {
max_failed_login_attempts_per_source_overrides?: {
[key: string]: number;
} | null;
/**
* Max File Downloads Per Minute
* @description max GET /v1/files/{file_id}/content calls per key (per user for JWT callers) per file per minute. A key's metadata can override it and a team's metadata adds a shared team cap, both set by a proxy admin. Unset means no limit
*/
max_file_downloads_per_minute?: number | null;
/**
* Max File Size Mb
* @description max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider