mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
12dcce8db2
commit
402fa63366
25 changed files with 3134 additions and 51 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
250
litellm/proxy/openai_files_endpoints/file_usage_caps.py
Normal file
250
litellm/proxy/openai_files_endpoints/file_usage_caps.py
Normal 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.",
|
||||
)
|
||||
|
|
@ -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)),
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
353
tests/integration/management/_batch_file_caps.py
Normal file
353
tests/integration/management/_batch_file_caps.py
Normal 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},
|
||||
)
|
||||
928
tests/integration/management/test_batch_file_usage_caps.py
Normal file
928
tests/integration/management/test_batch_file_usage_caps.py
Normal 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]
|
||||
411
tests/integration/management/test_batch_file_usage_caps_chaos.py
Normal file
411
tests/integration/management/test_batch_file_usage_caps_chaos.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
239
tests/unit/proxy/openai_files_endpoint/test_file_usage_caps.py
Normal file
239
tests/unit/proxy/openai_files_endpoint/test_file_usage_caps.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
15
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
15
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue