mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 4ee2d9a9de into d232b008bc
This commit is contained in:
commit
0630a95c36
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