diff --git a/litellm/constants.py b/litellm/constants.py index 33c93194ff2..a268f00775a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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. diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 78d57b67656..c21c88c3505 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 31e7483b891..829e792cf38 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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." }, ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 53a5b8de30a..5054ceb92d5 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 11810f4bd52..2ee29f894b5 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 133df8f9b3d..0266b452b61 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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, diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py index 21c0f03605b..af62ea62a1a 100644 --- a/litellm/proxy/management_helpers/bulk_user_creation.py +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/batch_file_validation.py b/litellm/proxy/openai_files_endpoints/batch_file_validation.py index fcd3be56ae7..d49aaaeba65 100644 --- a/litellm/proxy/openai_files_endpoints/batch_file_validation.py +++ b/litellm/proxy/openai_files_endpoints/batch_file_validation.py @@ -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=( diff --git a/litellm/proxy/openai_files_endpoints/file_usage_caps.py b/litellm/proxy/openai_files_endpoints/file_usage_caps.py new file mode 100644 index 00000000000..5de5bbce27e --- /dev/null +++ b/litellm/proxy/openai_files_endpoints/file_usage_caps.py @@ -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.", + ) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index dfec77a6c16..e51d912dcfd 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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)), diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d3d12df9feb..06bfb467b52 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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", diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5c8edf34917..d90ff1581e2 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/tests/integration/management/_batch_file_caps.py b/tests/integration/management/_batch_file_caps.py new file mode 100644 index 00000000000..e0625c48c42 --- /dev/null +++ b/tests/integration/management/_batch_file_caps.py @@ -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}, + ) diff --git a/tests/integration/management/test_batch_file_usage_caps.py b/tests/integration/management/test_batch_file_usage_caps.py new file mode 100644 index 00000000000..683425f3b1d --- /dev/null +++ b/tests/integration/management/test_batch_file_usage_caps.py @@ -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] diff --git a/tests/integration/management/test_batch_file_usage_caps_chaos.py b/tests/integration/management/test_batch_file_usage_caps_chaos.py new file mode 100644 index 00000000000..0ec3ab597e0 --- /dev/null +++ b/tests/integration/management/test_batch_file_usage_caps_chaos.py @@ -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 diff --git a/tests/integration/management/test_batch_limit_metadata_admin_only.py b/tests/integration/management/test_batch_limit_metadata_admin_only.py new file mode 100644 index 00000000000..163968e7b3e --- /dev/null +++ b/tests/integration/management/test_batch_limit_metadata_admin_only.py @@ -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= passed in."}} + assert config_entry(gateway, setting)["field_value"] is None diff --git a/tests/integration/observability/test_file_usage_counter_span.py b/tests/integration/observability/test_file_usage_counter_span.py new file mode 100644 index 00000000000..3d4ab9fab2c --- /dev/null +++ b/tests/integration/observability/test_file_usage_counter_span.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 799fe59147c..7150ba892df 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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` diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 15bf4f31445..ce8f05e879e 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -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( diff --git a/tests/unit/proxy/management_helpers/test_bulk_user_creation.py b/tests/unit/proxy/management_helpers/test_bulk_user_creation.py index b5349fc2387..cf1ac61c93a 100644 --- a/tests/unit/proxy/management_helpers/test_bulk_user_creation.py +++ b/tests/unit/proxy/management_helpers/test_bulk_user_creation.py @@ -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() diff --git a/tests/unit/proxy/openai_files_endpoint/test_file_usage_caps.py b/tests/unit/proxy/openai_files_endpoint/test_file_usage_caps.py new file mode 100644 index 00000000000..7cf2505dd38 --- /dev/null +++ b/tests/unit/proxy/openai_files_endpoint/test_file_usage_caps.py @@ -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 diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_batch_file_validation.py b/tests/unit/proxy/openai_files_endpoint/test_files_batch_file_validation.py index 3a73c39e177..b55ae1500be 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_batch_file_validation.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_batch_file_validation.py @@ -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) diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py index 84cf4ea7c32..25c62fc6dcc 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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) diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index 70c5a153647..c5a97fa4501 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 5c39bbd2cae..e4d9111ab7a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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