diff --git a/.github/actions/cache-prisma-binaries/action.yml b/.github/actions/cache-prisma-binaries/action.yml new file mode 100644 index 00000000000..68615e94c08 --- /dev/null +++ b/.github/actions/cache-prisma-binaries/action.yml @@ -0,0 +1,40 @@ +name: "Cache Prisma binaries" +description: >- + Cache the Prisma CLI and engine binaries that `prisma generate` downloads, so + only the first job on a given prisma-client-py version pays for the download. + + prisma-client-py shells out to `npm install prisma@` whenever its + binary cache directory has no CLI entrypoint, which pulls ~85 MB of query and + schema engines over the network. That normally takes a few seconds, but it is + unbounded: one shard of a proxy-db run took 5m18s on that single step versus + 3.8s on its eleven siblings, which pushed the job past its timeout and got a + fully passing test run cancelled. + + Callers must not set PRISMA_BINARY_CACHE_DIR. The prisma-client-py default + (~/.cache/prisma-python/binaries//) is already + keyed by both versions, so a cache entry can never be served to a run that + expects different binaries. + +runs: + using: composite + steps: + - name: Resolve prisma-client-py version + id: version + shell: bash + run: | + version="$(grep -A1 '^name = "prisma"$' uv.lock | sed -n 's/^version = "\(.*\)"$/\1/p' | head -1)" + if [ -z "${version}" ]; then + echo "could not resolve the prisma package version from uv.lock" >&2 + exit 1 + fi + echo "version=${version}" >> "$GITHUB_OUTPUT" + + - name: Restore Prisma binaries + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + with: + # ~/.cache/prisma-python holds the npm install tree prisma-client-py + # drives; ~/.cache/prisma is where @prisma/engines stages its downloads. + path: | + ~/.cache/prisma-python + ~/.cache/prisma + key: ${{ runner.os }}-prisma-binaries-${{ steps.version.outputs.version }} diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index cee93bde7f2..58208988fca 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -18,10 +18,25 @@ on: type: number default: 2 timeout-minutes: - description: "Job timeout in minutes" + description: >- + Timeout for the test step alone. Setup (checkout, dependency install, + Prisma client generation) gets its own allowance on top, so a slow + runner or a cold binary download can never cancel passing tests. required: false type: number default: 20 + job-timeout-minutes: + description: >- + Backstop for the whole job. Keep it >= `timeout-minutes` plus 35: 30 for + the per-step ceilings on the setup steps below, and 5 for the runner + overhead the job clock charges but no step owns (job init, step + transitions, post-job cleanup). That headroom is what makes the test + budget a floor rather than a hope, since setup cannot overrun into it + without failing its own step first. GitHub expressions have no + arithmetic, so the sum is passed in rather than computed. + required: false + type: number + default: 55 max-failures: description: "Stop after this many failures" required: false @@ -44,30 +59,35 @@ jobs: run: name: Run tests runs-on: ubuntu-latest - timeout-minutes: ${{ inputs.timeout-minutes }} + timeout-minutes: ${{ inputs.job-timeout-minutes }} outputs: decision: ${{ steps.changes.outputs.decision }} steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + timeout-minutes: 3 with: persist-credentials: false - name: Detect backend-relevant changes id: changes + timeout-minutes: 2 uses: ./.github/actions/detect-backend-changes - name: Set up Python + timeout-minutes: 3 uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 with: python-version: "3.12" - name: Set up uv + timeout-minutes: 3 uses: ./.github/actions/setup-uv-with-retries with: version: "0.10.9" - name: Cache uv dependencies + timeout-minutes: 5 uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 with: path: | @@ -79,18 +99,24 @@ jobs: - name: Install dependencies if: steps.changes.outputs.decision != 'skip' + timeout-minutes: 8 run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml + - name: Cache Prisma binaries + if: steps.changes.outputs.decision != 'skip' + timeout-minutes: 3 + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client if: steps.changes.outputs.decision != 'skip' - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache + timeout-minutes: 3 run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - name: Run tests if: steps.changes.outputs.decision != 'skip' + timeout-minutes: ${{ inputs.timeout-minutes }} env: TEST_PATH: ${{ inputs.test-path }} MAX_FAILURES: ${{ inputs.max-failures }} diff --git a/.github/workflows/check-ui-api-types.yml b/.github/workflows/check-ui-api-types.yml index 02543d67a82..dbd663a2efa 100644 --- a/.github/workflows/check-ui-api-types.yml +++ b/.github/workflows/check-ui-api-types.yml @@ -71,10 +71,12 @@ jobs: if: steps.changes.outputs.relevant == 'true' run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router + - name: Cache Prisma binaries + if: steps.changes.outputs.relevant == 'true' + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client if: steps.changes.outputs.relevant == 'true' - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - name: Set up Node.js diff --git a/.github/workflows/mutation-test.yml b/.github/workflows/mutation-test.yml index da4fe073a6a..68317d5dd12 100644 --- a/.github/workflows/mutation-test.yml +++ b/.github/workflows/mutation-test.yml @@ -57,9 +57,10 @@ jobs: run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml + - name: Cache Prisma binaries + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma diff --git a/.github/workflows/publish-basedpyright-base-counts.yml b/.github/workflows/publish-basedpyright-base-counts.yml index c85d30df0ce..71e196d8361 100644 --- a/.github/workflows/publish-basedpyright-base-counts.yml +++ b/.github/workflows/publish-basedpyright-base-counts.yml @@ -43,12 +43,13 @@ jobs: with: version: "0.10.9" + - name: Cache Prisma binaries + uses: ./.github/actions/cache-prisma-binaries + # The gate provisions its own measurement env (.venv-typecheck: a frozen # uv sync of its canonical dependency groups plus a generated Prisma # client), so no install step here can drift from what local runs measure. - name: Emit basedpyright counts for HEAD - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/basedpyright-counts" counts_file=$(ls "$RUNNER_TEMP"/basedpyright-counts/basedpyright-counts-*.json) diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index fab05fc2bbb..8f62837d29a 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -65,6 +65,12 @@ jobs: - name: check_provider_folders_documented run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py + - name: check_prisma_binary_cache + run: uv run --no-sync python ./tests/code_coverage_tests/check_prisma_binary_cache.py + + - name: check_workflow_startup_safety + run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py + - name: router_code_coverage run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 3db3fb07a94..280ec476cdf 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -71,12 +71,13 @@ jobs: run: | uv sync --frozen --group proxy-dev --group e2e-dev + - name: Cache Prisma binaries + uses: ./.github/actions/cache-prisma-binaries + # basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma) # only after `prisma generate` writes prisma/client.py et al. Without this the # DB wrappers typed against the generated client would degrade to Unknown. - name: Generate Prisma client - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma @@ -119,7 +120,6 @@ jobs: - name: Check basedpyright budget (delta vs base) env: GH_TOKEN: ${{ github.token }} - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA" diff --git a/.github/workflows/test-terraform-provider.yml b/.github/workflows/test-terraform-provider.yml index 058a2538c15..7ea22825f4f 100644 --- a/.github/workflows/test-terraform-provider.yml +++ b/.github/workflows/test-terraform-provider.yml @@ -92,9 +92,10 @@ jobs: run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router + - name: Cache Prisma binaries + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml index 50589cb5926..c93779c177f 100644 --- a/.github/workflows/test-unit-documentation.yml +++ b/.github/workflows/test-unit-documentation.yml @@ -65,10 +65,12 @@ jobs: run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router + - name: Cache Prisma binaries + if: steps.changes.outputs.decision != 'skip' + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client if: steps.changes.outputs.decision != 'skip' - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 60d2e471862..df212a85885 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -28,6 +28,10 @@ concurrency: # Most of a shard's time is pytest plugin load + xdist worker imports + # pytest-cov instrumentation, not the tests themselves. Keeping per-shard # work low and matching worker count to runner cores is what controls it. +# * `timeout` bounds the pytest step only. Checkout, dependency install, and +# Prisma client generation draw on a separate allowance in the base +# workflow, so slow setup shows up as a slow job rather than as a +# cancelled shard whose tests were passing. # * workers: 4 matches the 4-core ubuntu-latest runner. -n 8 on 4 cores # oversubscribes 2x and workers fight for CPU during their cold-start # imports (measured ~441% CPU for -n 8 locally, i.e. ~55% effective). diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index 2ea3c521e8b..64b92f7d847 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -76,4 +76,5 @@ jobs: workers: 4 reruns: 2 timeout-minutes: 60 + job-timeout-minutes: 95 artifact-name: proxy-server diff --git a/.github/workflows/test-unit-proxy-legacy.yml b/.github/workflows/test-unit-proxy-legacy.yml index 49aa5f9f51d..e8ca36fb30d 100644 --- a/.github/workflows/test-unit-proxy-legacy.yml +++ b/.github/workflows/test-unit-proxy-legacy.yml @@ -82,10 +82,12 @@ jobs: run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router + - name: Cache Prisma binaries + if: steps.changes.outputs.decision != 'skip' + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client if: steps.changes.outputs.decision != 'skip' - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma diff --git a/.github/workflows/weekly_load_anomaly.yml b/.github/workflows/weekly_load_anomaly.yml index 4c2103f026d..2dffc889d0e 100644 --- a/.github/workflows/weekly_load_anomaly.yml +++ b/.github/workflows/weekly_load_anomaly.yml @@ -51,9 +51,10 @@ jobs: run: | .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy + - name: Cache Prisma binaries + uses: ./.github/actions/cache-prisma-binaries + - name: Generate Prisma client - env: - PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache run: | uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 65d3c239253..96b689aed74 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45004 + "limit": 44996 }, "reportUnknownLambdaType": { "limit": 113 }, "reportUnknownMemberType": { - "limit": 39649 + "limit": 39643 }, "reportUnknownParameterType": { "limit": 20132 }, "reportUnknownVariableType": { - "limit": 31156 + "limit": 31153 }, "reportUnnecessaryCast": { "limit": 118 diff --git a/litellm/constants.py b/litellm/constants.py index f2ac96162eb..87d6fa1a744 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1493,6 +1493,8 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)) +RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500"))) +RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", "100"))) PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0e0f1558bb5..c8ca9dcf57e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,7 +1,7 @@ import enum import json import os -from collections.abc import Callable +from collections.abc import Callable, Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal @@ -11,6 +11,7 @@ from pydantic import ( ConfigDict, Field, Json, + PositiveInt, field_validator, model_validator, ) @@ -1102,6 +1103,8 @@ class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase): class KeyRequestBase(GenerateRequestBase): key: str | None = None + default_estimated_output_tokens: PositiveInt | None = None + default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None budget_id: str | None = None tags: list[str] | None = None disable_global_guardrails: bool | None = None @@ -1819,6 +1822,8 @@ class NewTeamRequest(TeamBase): ) model_tpm_limit: dict[str, int] | None = None + default_estimated_output_tokens: PositiveInt | None = None + default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None mcp_rpm_limit: dict[str, int] | None = None team_member_budget: float | None = None # allow user to set a budget for all team members team_member_rpm_limit: int | None = None # allow user to set RPM limit for all team members @@ -1883,6 +1888,8 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): prompts: list[str] | None = None model_rpm_limit: dict[str, int] | None = None model_tpm_limit: dict[str, int] | None = None + default_estimated_output_tokens: PositiveInt | None = None + default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None mcp_rpm_limit: dict[str, int] | None = None allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None enforced_batch_output_expires_after: dict | None = None @@ -4103,6 +4110,8 @@ class PassThroughEndpointLoggingTypedDict(TypedDict): LiteLLM_ManagementEndpoint_MetadataFields: Final = [ "model_rpm_limit", "model_tpm_limit", + "default_estimated_output_tokens", + "default_estimated_output_tokens_per_model", "mcp_rpm_limit", "tag_rpm_limit", "rpm_limit_type", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index b4e634b5eb1..0bfe4b685e1 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1,12 +1,13 @@ import os import re import sys -from collections.abc import Iterator, Mapping +from collections.abc import Collection, Iterator, Mapping from functools import lru_cache from logging import Logger -from typing import Any, Final +from typing import Any, Final, Protocol from fastapi import HTTPException, Request, status +from pydantic import PositiveInt, TypeAdapter, ValidationError import litellm from litellm import Router, provider_list @@ -999,6 +1000,167 @@ def get_key_model_tpm_limit( return None +ESTIMATED_OUTPUT_TOKENS_FIELD: Final = "default_estimated_output_tokens" +ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD: Final = "default_estimated_output_tokens_per_model" +ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS: Final = frozenset( + {ESTIMATED_OUTPUT_TOKENS_FIELD, ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD} +) + +_ESTIMATED_OUTPUT_TOKENS_ADAPTER: Final = TypeAdapter(PositiveInt) +_ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER: Final = TypeAdapter(Mapping[str, PositiveInt]) + + +def _validated_output_token_estimate(raw: object) -> int | None: + """Coerce one declared estimate to a positive int, or ignore it.""" + if raw is None: + return None + try: + return _ESTIMATED_OUTPUT_TOKENS_ADAPTER.validate_python(raw) + except ValidationError as validation_error: + verbose_proxy_logger.warning( + "Ignoring malformed %s in metadata: %s", + ESTIMATED_OUTPUT_TOKENS_FIELD, + validation_error, + ) + return None + + +def _validated_output_token_estimates_per_model(raw: object) -> Mapping[str, int] | None: + """Coerce a declared per-model estimate map, or ignore it.""" + if raw is None: + return None + try: + return _ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER.validate_python(raw) + except ValidationError as validation_error: + verbose_proxy_logger.warning( + "Ignoring malformed %s in metadata: %s", + ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD, + validation_error, + ) + return None + + +def _estimated_output_tokens_from_metadata( + metadata: Mapping[str, Any] | None, + model_name: str | None, +) -> int | None: + """Resolve the per-model, then global, estimate out of one metadata blob. + + The two fields are validated independently so a malformed per-model map + cannot discard a valid global estimate, or the other way round. + """ + if not metadata or ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS.isdisjoint(metadata): + return None + + if model_name is not None: + per_model: Final = _validated_output_token_estimates_per_model( + metadata.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD) + ) + per_model_estimate: Final = per_model.get(model_name) if per_model is not None else None + if per_model_estimate is not None: + return per_model_estimate + + return _validated_output_token_estimate(metadata.get(ESTIMATED_OUTPUT_TOKENS_FIELD)) + + +def get_estimated_output_tokens( + user_api_key_dict: UserAPIKeyAuth, + model_name: str | None = None, +) -> int | None: + """Resolve the operator-declared output-token estimate for TPM reservation. + + Priority order (returns first found): + 1. Key metadata ``default_estimated_output_tokens_per_model[model_name]`` + 2. Key metadata ``default_estimated_output_tokens`` + 3. Team metadata ``default_estimated_output_tokens_per_model[model_name]`` + 4. Team metadata ``default_estimated_output_tokens`` + + Returns ``None`` when nothing is configured, which leaves the static + heuristic floor in place. + """ + key_estimate: Final = _estimated_output_tokens_from_metadata(user_api_key_dict.metadata, model_name) + if key_estimate is not None: + return key_estimate + return _estimated_output_tokens_from_metadata(user_api_key_dict.team_metadata, model_name) + + +class OutputTokenEstimateRequest(Protocol): + """The shape of any management request that can carry an output-token estimate. + + Read-only members: the gate inspects a request, it never writes one back. + """ + + @property + def metadata(self) -> Mapping[str, object] | None: ... + + @property + def default_estimated_output_tokens(self) -> int | None: ... + + @property + def default_estimated_output_tokens_per_model(self) -> Mapping[str, int] | None: ... + + @property + def model_fields_set(self) -> Collection[str]: ... + + +def _requested_output_token_estimates( + data: OutputTokenEstimateRequest, + existing_metadata: Mapping[str, object], +) -> tuple[object, object]: + """The output-token estimates this request would leave stored on the entity. + + Mirrors how the management endpoints merge metadata: a supplied ``metadata`` + replaces the stored blob wholesale, an omitted one preserves it, and the + dedicated top-level fields overlay whatever survives. Both sources are read + because the same declaration reaches the same stored field either way. + """ + base: Final[Mapping[str, object]] = ( + (data.metadata or {}) if "metadata" in data.model_fields_set else existing_metadata + ) + return ( + data.default_estimated_output_tokens + if data.default_estimated_output_tokens is not None + else base.get(ESTIMATED_OUTPUT_TOKENS_FIELD), + data.default_estimated_output_tokens_per_model + if data.default_estimated_output_tokens_per_model is not None + else base.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD), + ) + + +def enforce_output_token_estimates_are_admin_only( + data: OutputTokenEstimateRequest, + existing_metadata: Mapping[str, object] | None, + user_api_key_dict: UserAPIKeyAuth, + entity: Literal["key", "team"], +) -> None: + """Only a proxy admin may change what a key or team declares its models emit. + + That declaration is what the TPM limiter reserves for a request omitting + ``max_tokens``, so lowering or clearing it under-reserves against every + window the request is charged against, including the team and organization + ones the writer may not own. A key's metadata is writable by its holder and + a team's by its team admin, so neither is a trustworthy source for a value + that weakens a limit set above them. Gated on the resulting value rather + than on presence, so a form resending the stored declaration stays a no-op. + """ + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: + return + stored: Final[Mapping[str, object]] = existing_metadata or {} + if _requested_output_token_estimates(data, stored) == ( + stored.get(ESTIMATED_OUTPUT_TOKENS_FIELD), + stored.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD), + ): + return + raise HTTPException( + status_code=403, + detail={ + "error": f"Only proxy admins can set {ESTIMATED_OUTPUT_TOKENS_FIELD} or " + f"{ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD} on a {entity}. They decide how many output tokens " + "the rate limiter reserves for a request that omits max_tokens." + }, + ) + + def get_model_rate_limit_from_metadata( user_api_key_dict: UserAPIKeyAuth, metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"], diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 8830970f96f..bf760a92d88 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -1,14 +1,21 @@ import asyncio import json import time -from collections.abc import Callable, Sequence +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass from datetime import datetime, timezone -from typing import Final, Literal, Protocol, TypeVar +from types import MappingProxyType +from typing import Final, Literal, Protocol, TypeVar, assert_never import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache -from litellm.constants import GLOBAL_PROXY_SPEND_CACHE_KEY, LITELLM_PROXY_BUDGET_NAME +from litellm.constants import ( + GLOBAL_PROXY_SPEND_CACHE_KEY, + LITELLM_PROXY_BUDGET_NAME, + RESET_BUDGET_JOB_BATCH_SIZE, + RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN, +) from litellm.proxy._types import ( LiteLLM_BudgetTableFull, LiteLLM_EndUserTable, @@ -30,7 +37,10 @@ from litellm.repositories.table_repositories import ( TeamMembershipRepository, ) from litellm.repositories.team_repository import TeamRepository -from litellm.repositories.unit_of_work import spend_reset_unit_of_work +from litellm.repositories.unit_of_work import ( + budget_cascade_unit_of_work, + spend_reset_unit_of_work, +) from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) @@ -38,6 +48,9 @@ from litellm.types.services import ServiceTypes _RowT = TypeVar("_RowT") +_LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_duration": None, "spend": {"gt": 0}}) +_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}}) + class _TeamMembershipRow(Protocol): @property @@ -62,39 +75,130 @@ class _TagRow(Protocol): def tag_name(self) -> str: ... +class _EndUserRow(Protocol): + @property + def user_id(self) -> str: ... + + def _team_membership_counter_key(row: _TeamMembershipRow) -> str: return f"spend:team_member:{row.user_id}:{row.team_id}" -def _team_membership_cache_key(row: _TeamMembershipRow) -> str: - return f"{row.team_id}_{row.user_id}" +def _team_membership_cache_keys(row: _TeamMembershipRow) -> tuple[str, ...]: + return (f"{row.team_id}_{row.user_id}",) def _key_counter_key(row: _KeyRow) -> str: return f"spend:key:{row.token}" -def _key_cache_key(row: _KeyRow) -> str: - return row.token +def _key_cache_keys(row: _KeyRow) -> tuple[str, ...]: + return (row.token,) def _org_counter_key(row: _OrgRow) -> str: return f"spend:org:{row.organization_id}" -def _org_cache_keys(row: _OrgRow) -> Sequence[str]: - return [ +def _org_cache_keys(row: _OrgRow) -> tuple[str, ...]: + return ( f"org_id:{row.organization_id}", f"org_id:{row.organization_id}:with_budget", - ] + ) def _tag_counter_key(row: _TagRow) -> str: return f"spend:tag:{row.tag_name}" -def _tag_cache_key(row: _TagRow) -> str: - return f"tag:{row.tag_name}" +def _tag_cache_keys(row: _TagRow) -> tuple[str, ...]: + return (f"tag:{row.tag_name}",) + + +def _budget_link_where( + budget_ids: Sequence[str], + extra: Mapping[str, object] = MappingProxyType({}), +) -> dict[str, object]: + return {"budget_id": {"in": list(budget_ids)}, **extra} + + +@dataclass(frozen=True, slots=True) +class _BudgetCascade: + """Everything one budget-tier reset touches, resolved before any write.""" + + budgets: tuple[LiteLLM_BudgetTableFull, ...] = () + budget_ids: tuple[str, ...] = () + budget_resets: tuple[tuple[str, datetime], ...] = () + endusers: tuple[_EndUserRow, ...] = () + counter_keys: tuple[str, ...] = () + cache_keys: tuple[str, ...] = () + + +@dataclass(frozen=True, slots=True) +class _BudgetCascadeCommitted: + cascade: _BudgetCascade + advanced: int + + +@dataclass(frozen=True, slots=True) +class _BudgetCascadeFailed: + cascade: _BudgetCascade + error: Exception + + +_EMPTY_CASCADE: Final = _BudgetCascade() + + +@dataclass(frozen=True, slots=True) +class _ChunkOutcome: + """One chunk of a reset phase: rows read, and rows whose new budget_reset_at + cleared the due cutoff. Anything else is still due and would come straight + back on the next fetch, so it is not progress.""" + + fetched: int + advanced: int + + +_NO_PROGRESS: Final = _ChunkOutcome(fetched=0, advanced=0) + + +def _as_utc(moment: datetime) -> datetime: + return moment if moment.tzinfo is not None else moment.replace(tzinfo=timezone.utc) + + +def _count_advanced(reset_ats: Iterable[object], cutoff: datetime) -> int: + """How many rows the write actually moved past the due cutoff. + + A budget_duration of "0s" (or one the parser cannot read) resolves to the + current time, so the row is written and stays due. Counting it as progress + would re-read the same chunk until the per-run cap on every tick. + """ + utc_cutoff: Final = _as_utc(cutoff) + return sum(1 for reset_at in reset_ats if isinstance(reset_at, datetime) and _as_utc(reset_at) > utc_cutoff) + + +def _phase_is_drained(outcome: _ChunkOutcome) -> bool: + """A short chunk means the due rows ran out. A full chunk that advanced + nothing would be re-read unchanged forever, so it ends the phase too and + those rows wait for the next tick.""" + return outcome.fetched < RESET_BUDGET_JOB_BATCH_SIZE or outcome.advanced == 0 + + +async def _run_phase_in_chunks(process_chunk: Callable[[], Awaitable[_ChunkOutcome]]) -> None: + """Drive one reset phase a chunk at a time, capped so a single run cannot + spin unbounded: leftovers are picked up by the next tick.""" + for _ in range(RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN): + if _phase_is_drained(await process_chunk()): + return + + +def _budget_cascade_event_metadata(cascade: _BudgetCascade) -> dict[str, object]: + return { + "num_budgets_found": len(cascade.budgets), + "budgets_found": json.dumps(cascade.budgets, indent=4, default=str), + "num_endusers_found": len(cascade.endusers), + "endusers_found": json.dumps(cascade.endusers, indent=4, default=str), + } class ResetBudgetJob: @@ -122,21 +226,14 @@ class ResetBudgetJob: Updates db """ - if self.prisma_client is not None: - ### RESET KEY BUDGET ### - await self.reset_budget_for_litellm_keys() + if self.prisma_client is None: + return - ### RESET USER BUDGET ### - await self.reset_budget_for_litellm_users() - - ## Reset Team Budget - await self.reset_budget_for_litellm_teams() - - ### RESET ENDUSER (Customer) BUDGET and corresponding Budget duration ### - await self.reset_budget_for_litellm_budget_table() - - ### RESET MULTI-WINDOW BUDGETS ### - await self.reset_budget_windows() + await self.reset_budget_for_litellm_keys() + await self.reset_budget_for_litellm_users() + await self.reset_budget_for_litellm_teams() + await self.reset_budget_for_litellm_budget_table() + await self.reset_budget_windows() @staticmethod async def _invalidate_spend_counter(counter_key: str) -> None: @@ -194,238 +291,195 @@ class ResetBudgetJob: e, ) - async def _cascade_reset_spend_for_budget_link( + async def _fetch_linked_rows( self, - budgets_to_reset: list[LiteLLM_BudgetTableFull], table: SpendLinkedTable[_RowT], - counter_key_fn: Callable[[_RowT], str], + where: Mapping[str, object], log_subject: str, - extra_where: dict[str, object] | None = None, - cache_key_fn: Callable[[_RowT], str | Sequence[str]] | None = None, - ): - """ - Generic cascade: zero spend on rows whose budget_id is in the reset set. + ) -> tuple[_RowT, ...]: + """Read the rows the cascade will zero, so their counters can be + invalidated once the transaction commits.""" + try: + return tuple(await table.find_many(where=where)) + except Exception as e: + verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e) + return () - ``cache_key_fn`` is optional: when provided, after the DB update each - matching row's entry or entries in ``user_api_key_cache`` are dropped so - cached spend cannot stay pinned above the zeroed DB row after a reset. + async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]: + linked: Final[Sequence[_EndUserRow] | None] = await self.prisma_client.get_data( + table_name="enduser", + query_type="find_all", + budget_id_list=list(budget_ids), + ) + if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids: + return tuple(linked or ()) + return (*(linked or ()), *await self._get_endusers_with_no_budget_id()) + + async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade: + """Resolve every row the expiring budget tiers gate, before any write. + + Keys carrying their own budget_duration are left out: they run on their + own schedule via reset_budget_for_litellm_keys(), so sweeping them here + would reset them twice. """ - budget_ids: Final = [b.budget_id for b in budgets_to_reset if b.budget_id is not None] + budget_ids: Final = tuple(b.budget_id for b in budgets_to_reset if b.budget_id is not None) if not budget_ids: + return _EMPTY_CASCADE + + team_memberships: Final[tuple[_TeamMembershipRow, ...]] = await self._fetch_linked_rows( + table=TeamMembershipRepository(self.prisma_client).table, + where=_budget_link_where(budget_ids), + log_subject="team memberships", + ) + keys: Final[tuple[_KeyRow, ...]] = await self._fetch_linked_rows( + table=VerificationTokenRepository(self.prisma_client).table, + where=_budget_link_where(budget_ids, _LINKED_KEYS_WHERE), + log_subject="keys", + ) + orgs: Final[tuple[_OrgRow, ...]] = await self._fetch_linked_rows( + table=OrganizationRepository(self.prisma_client).table, + where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE), + log_subject="orgs", + ) + tags: Final[tuple[_TagRow, ...]] = await self._fetch_linked_rows( + table=TagRepository(self.prisma_client).table, + where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE), + log_subject="tags", + ) + return _BudgetCascade( + budgets=tuple(budgets_to_reset), + budget_ids=budget_ids, + budget_resets=tuple( + ( + b.budget_id, + compute_budget_reset_at(budget_duration=b.budget_duration, settings=self.reset_settings), + ) + for b in budgets_to_reset + if b.budget_id is not None and b.budget_duration is not None + ), + endusers=await self._collect_endusers_to_reset(budget_ids), + counter_keys=( + *(_team_membership_counter_key(row) for row in team_memberships), + *(_key_counter_key(row) for row in keys), + *(_org_counter_key(row) for row in orgs), + *(_tag_counter_key(row) for row in tags), + ), + cache_keys=( + *(key for row in team_memberships for key in _team_membership_cache_keys(row)), + *(key for row in keys for key in _key_cache_keys(row)), + *(key for row in orgs for key in _org_cache_keys(row)), + *(key for row in tags for key in _tag_cache_keys(row)), + ), + ) + + async def _commit_budget_cascade(self, cascade: _BudgetCascade) -> None: + """Zero the gated spend and advance ``budget_reset_at`` in one transaction. + + Advancing the window on its own hides the tier from every later tick + while its dependents stay pinned at the cap for the whole window; + batching both means a mid-cascade failure persists nothing and the rows + stay due for the next run. + """ + if not cascade.budget_ids: return - where: Final[dict[str, object]] = {"budget_id": {"in": budget_ids}} - if extra_where: - where.update(extra_where) + enduser_ids: Final = tuple(row.user_id for row in cascade.endusers) + async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow: + uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids)) + uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE)) + uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE)) + uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE)) + if enduser_ids: + uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}}) + for budget_id, budget_reset_at in cascade.budget_resets: + uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at) - try: - rows: Sequence[_RowT] = await table.find_many(where=where) - except Exception as e: - rows = () - verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e) - - update_result: Final = await table.update_many(where=where, data={"spend": 0}) - - for row in rows: - await self._invalidate_spend_counter(counter_key_fn(row)) - if cache_key_fn is not None: - cache_keys = cache_key_fn(row) - if isinstance(cache_keys, str): - cache_keys = [cache_keys] - for cache_key in cache_keys: - await self._invalidate_user_api_key_cache_entry(cache_key) - - return update_result - - async def reset_budget_for_litellm_team_members(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): - """ - Resets the budget for all LiteLLM Team Members if their budget has expired - """ - return await self._cascade_reset_spend_for_budget_link( - budgets_to_reset=budgets_to_reset, - table=TeamMembershipRepository(self.prisma_client).table, - counter_key_fn=_team_membership_counter_key, - log_subject="team memberships", - cache_key_fn=_team_membership_cache_key, - ) - - async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): - """ - Resets the spend for keys linked to budget tiers that are being reset. - - Excludes keys with their own budget_duration; those are reset by - reset_budget_for_litellm_keys() to avoid double-resetting. - """ - return await self._cascade_reset_spend_for_budget_link( - budgets_to_reset=budgets_to_reset, - table=VerificationTokenRepository(self.prisma_client).table, - counter_key_fn=_key_counter_key, - log_subject="keys", - extra_where={"budget_duration": None, "spend": {"gt": 0}}, - cache_key_fn=_key_cache_key, - ) - - async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): - """ - Resets the spend for orgs linked to budget tiers that are being reset. - """ - return await self._cascade_reset_spend_for_budget_link( - budgets_to_reset=budgets_to_reset, - table=OrganizationRepository(self.prisma_client).table, - counter_key_fn=_org_counter_key, - log_subject="orgs", - extra_where={"spend": {"gt": 0}}, - cache_key_fn=_org_cache_keys, - ) - - async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]): - """ - Resets the spend for tags linked to budget tiers that are being reset. - - Also drops each tag's ``user_api_key_cache`` entry so the next - ``_tag_max_budget_check`` reloads the zeroed row from the DB. - ``SpendCounterReseed.from_db`` intentionally returns ``None`` for - tags, so the budget check falls back to the cached - ``LiteLLM_TagTable.spend`` once the spend counter expires; without - this invalidation, that stale ``.spend`` keeps the tag over-budget - indefinitely. - """ - return await self._cascade_reset_spend_for_budget_link( - budgets_to_reset=budgets_to_reset, - table=TagRepository(self.prisma_client).table, - counter_key_fn=_tag_counter_key, - log_subject="tags", - extra_where={"spend": {"gt": 0}}, - cache_key_fn=_tag_cache_key, - ) - - async def reset_budget_for_litellm_budget_table(self): - """ - Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired - The corresponding Budget duration is also updated. - """ + async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None: + for counter_key in cascade.counter_keys: + await self._invalidate_spend_counter(counter_key) + for cache_key in cascade.cache_keys: + await self._invalidate_user_api_key_cache_entry(cache_key) + async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed: now: Final = datetime.now(timezone.utc) - start_time: Final = time.time() - endusers_to_reset: list[LiteLLM_EndUserTable] | None = None - budgets_to_reset: list[LiteLLM_BudgetTableFull] | None = None - updated_endusers: Final[list[LiteLLM_EndUserTable]] = [] - failed_endusers: Final = [] try: - budgets_to_reset = await self.prisma_client.get_data( - table_name="budget", query_type="find_all", reset_at=now - ) - - if budgets_to_reset is not None and len(budgets_to_reset) > 0: - for budget in budgets_to_reset: - budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings) - - await self.prisma_client.update_data( - query_type="update_many", - data_list=budgets_to_reset, - table_name="budget", - ) - - budget_ids_to_reset = [budget.budget_id for budget in budgets_to_reset if budget.budget_id is not None] - - endusers_to_reset = await self.prisma_client.get_data( - table_name="enduser", - query_type="find_all", - budget_id_list=budget_ids_to_reset, - ) - - # Also reset end users with no budget_id (NULL) who use the - # default budget via litellm.max_end_user_budget_id. These - # users are enforced in-memory but never had budget_id - # persisted, so the query above misses them. - if litellm.max_end_user_budget_id is not None and litellm.max_end_user_budget_id in budget_ids_to_reset: - default_budget_endusers: Final = await self._get_endusers_with_no_budget_id() - if default_budget_endusers: - if endusers_to_reset is None: - endusers_to_reset = default_budget_endusers - else: - endusers_to_reset.extend(default_budget_endusers) - - await self.reset_budget_for_litellm_team_members(budgets_to_reset=budgets_to_reset) - - await self.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset) - - await self.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=budgets_to_reset) - - await self.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=budgets_to_reset) - - if endusers_to_reset is not None and len(endusers_to_reset) > 0: - for enduser in endusers_to_reset: - try: - updated_enduser = await ResetBudgetJob._reset_budget_for_enduser(enduser=enduser) - if updated_enduser is not None: - updated_endusers.append(updated_enduser) - else: - failed_endusers.append( - { - "enduser": enduser, - "error": "Returned None without exception", - } - ) - except Exception as e: - failed_endusers.append({"enduser": enduser, "error": str(e)}) - verbose_proxy_logger.exception("Failed to reset budget for enduser: %s", enduser) - - verbose_proxy_logger.debug( - "Updated users %s", - json.dumps(updated_endusers, indent=4, default=str), - ) - - await self.prisma_client.update_data( - query_type="update_many", - data_list=updated_endusers, - table_name="enduser", - ) - - end_time = time.time() - if len(failed_endusers) > 0: # If any endusers failed to reset - raise Exception( - f"Failed to reset {len(failed_endusers)} endusers: {json.dumps(failed_endusers, default=str)}" - ) - - asyncio.create_task( - self.proxy_logging_obj.service_logging_obj.async_service_success_hook( - service=ServiceTypes.RESET_BUDGET_JOB, - duration=end_time - start_time, - call_type="reset_budget_budget_table", - start_time=start_time, - end_time=end_time, - event_metadata={ - "num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0), - "budgets_found": json.dumps(budgets_to_reset, indent=4, default=str), - "num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0), - "endusers_found": json.dumps(endusers_to_reset, indent=4, default=str), - "num_endusers_updated": len(updated_endusers), - "endusers_updated": json.dumps(updated_endusers, indent=4, default=str), - "num_endusers_failed": len(failed_endusers), - "endusers_failed": json.dumps(failed_endusers, indent=4, default=str), - }, - ) + budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self.prisma_client.get_data( + table_name="budget", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, ) + cascade: Final = await self._collect_budget_cascade(budgets_to_reset or ()) except Exception as e: - end_time = time.time() - asyncio.create_task( - self.proxy_logging_obj.service_logging_obj.async_service_failure_hook( - service=ServiceTypes.RESET_BUDGET_JOB, - duration=end_time - start_time, - error=e, - call_type="reset_budget_endusers", - start_time=start_time, - end_time=end_time, - event_metadata={ - "num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0), - "budgets_found": json.dumps(budgets_to_reset, indent=4, default=str), - "num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0), - "endusers_found": json.dumps(endusers_to_reset, indent=4, default=str), - }, + return _BudgetCascadeFailed(cascade=_EMPTY_CASCADE, error=e) + + try: + await self._commit_budget_cascade(cascade) + except Exception as e: + return _BudgetCascadeFailed(cascade=cascade, error=e) + + await self._invalidate_budget_cascade_caches(cascade) + return _BudgetCascadeCommitted( + cascade=cascade, + advanced=_count_advanced( + (reset_at for _, reset_at in cascade.budget_resets), + cutoff=datetime.now(timezone.utc), + ), + ) + + async def reset_budget_for_litellm_budget_table(self) -> None: + """ + Resets the spend a budget tier gates (end users, team members, keys, + orgs, tags) and advances the tier's budget_reset_at, atomically. + + Caches are invalidated only after the transaction commits, so a failed + run cannot leave a zeroed counter in front of an un-reset DB row. + """ + await _run_phase_in_chunks(self._reset_budget_for_litellm_budget_table_chunk) + + async def _reset_budget_for_litellm_budget_table_chunk(self) -> _ChunkOutcome: + start_time: Final = time.time() + outcome: Final = await self._reset_expired_budget_cascade() + end_time: Final = time.time() + + match outcome: + case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced): + asyncio.create_task( + self.proxy_logging_obj.service_logging_obj.async_service_success_hook( + service=ServiceTypes.RESET_BUDGET_JOB, + duration=end_time - start_time, + call_type="reset_budget_budget_table", + start_time=start_time, + end_time=end_time, + event_metadata={ + **_budget_cascade_event_metadata(cascade), + "num_endusers_updated": len(cascade.endusers), + "num_endusers_failed": 0, + }, + ) ) - ) - verbose_proxy_logger.exception("Failed to reset budget for endusers: %s", e) + return _ChunkOutcome(fetched=len(cascade.budgets), advanced=advanced) + case _BudgetCascadeFailed(cascade=cascade, error=error): + verbose_proxy_logger.exception( + "Failed to reset the budget table cascade (team member, enduser, org and tag spend, plus " + "budget_reset_at); nothing was committed and the budgets stay due for the next run: %s", + error, + exc_info=error, + ) + asyncio.create_task( + self.proxy_logging_obj.service_logging_obj.async_service_failure_hook( + service=ServiceTypes.RESET_BUDGET_JOB, + duration=end_time - start_time, + error=error, + call_type="reset_budget_endusers", + start_time=start_time, + end_time=end_time, + event_metadata=_budget_cascade_event_metadata(cascade), + ) + ) + return _NO_PROGRESS + case _: + assert_never(outcome) async def _get_endusers_with_no_budget_id( self, @@ -486,18 +540,50 @@ class ResetBudgetJob: for t in updated_teams: uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at) - async def reset_budget_for_litellm_keys(self): + def _emit_phase_failure( + self, + call_type: str, + error: Exception, + start_time: float, + end_time: float, + event_metadata: dict[str, object], + ) -> None: + """Report rows that could not be reset without failing the chunk: the + rows that did reset are already committed, and raising here would cost + the phase every remaining chunk this tick. + """ + verbose_proxy_logger.error("%s: %s", call_type, error) + asyncio.create_task( + self.proxy_logging_obj.service_logging_obj.async_service_failure_hook( + service=ServiceTypes.RESET_BUDGET_JOB, + duration=end_time - start_time, + error=error, + call_type=call_type, + start_time=start_time, + end_time=end_time, + event_metadata=event_metadata, + ) + ) + + async def reset_budget_for_litellm_keys(self) -> None: """ Resets the budget for all the litellm keys Catches Exceptions and logs them """ + await _run_phase_in_chunks(self._reset_budget_for_litellm_keys_chunk) + + async def _reset_budget_for_litellm_keys_chunk(self) -> _ChunkOutcome: now: Final = datetime.utcnow() start_time: Final = time.time() keys_to_reset: list[LiteLLM_VerificationToken] | None = None try: keys_to_reset = await self.prisma_client.get_data( - table_name="key", query_type="find_all", expires=now, reset_at=now + table_name="key", + query_type="find_all", + expires=now, + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, ) verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str)) updated_keys: Final[list[LiteLLM_VerificationToken]] = [] @@ -528,8 +614,25 @@ class ResetBudgetJob: await self._invalidate_spend_counter(f"spend:key:{token}") end_time = time.time() - if len(failed_keys) > 0: # If any keys failed to reset - raise Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}") + outcome: Final = _ChunkOutcome( + fetched=len(keys_to_reset) if keys_to_reset else 0, + advanced=_count_advanced( + (k.budget_reset_at for k in updated_keys), + cutoff=datetime.now(timezone.utc), + ), + ) + if len(failed_keys) > 0: + self._emit_phase_failure( + call_type="reset_budget_keys", + error=Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}"), + start_time=start_time, + end_time=end_time, + event_metadata={ + "num_keys_found": len(keys_to_reset) if keys_to_reset else 0, + "keys_found": json.dumps(keys_to_reset, indent=4, default=str), + }, + ) + return outcome asyncio.create_task( self.proxy_logging_obj.service_logging_obj.async_service_success_hook( @@ -565,16 +668,27 @@ class ResetBudgetJob: ) ) verbose_proxy_logger.exception("Failed to reset budget for keys: %s", e) + return _NO_PROGRESS + else: + return outcome - async def reset_budget_for_litellm_users(self): + async def reset_budget_for_litellm_users(self) -> None: """ Resets the budget for all LiteLLM Internal Users if their budget has expired """ + await _run_phase_in_chunks(self._reset_budget_for_litellm_users_chunk) + + async def _reset_budget_for_litellm_users_chunk(self) -> _ChunkOutcome: now: Final = datetime.utcnow() start_time: Final = time.time() users_to_reset: list[LiteLLM_UserTable] | None = None try: - users_to_reset = await self.prisma_client.get_data(table_name="user", query_type="find_all", reset_at=now) + users_to_reset = await self.prisma_client.get_data( + table_name="user", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ) updated_users: Final[list[LiteLLM_UserTable]] = [] failed_users: Final = [] if users_to_reset is not None and len(users_to_reset) > 0: @@ -609,8 +723,27 @@ class ResetBudgetJob: await self._invalidate_global_proxy_spend_cache() end_time = time.time() - if len(failed_users) > 0: # If any users failed to reset - raise Exception(f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}") + outcome: Final = _ChunkOutcome( + fetched=len(users_to_reset) if users_to_reset else 0, + advanced=_count_advanced( + (u.budget_reset_at for u in updated_users), + cutoff=datetime.now(timezone.utc), + ), + ) + if len(failed_users) > 0: + self._emit_phase_failure( + call_type="reset_budget_users", + error=Exception( + f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}" + ), + start_time=start_time, + end_time=end_time, + event_metadata={ + "num_users_found": len(users_to_reset) if users_to_reset else 0, + "users_found": json.dumps(users_to_reset, indent=4, default=str), + }, + ) + return outcome asyncio.create_task( self.proxy_logging_obj.service_logging_obj.async_service_success_hook( @@ -646,16 +779,27 @@ class ResetBudgetJob: ) ) verbose_proxy_logger.exception("Failed to reset budget for users: %s", e) + return _NO_PROGRESS + else: + return outcome - async def reset_budget_for_litellm_teams(self): + async def reset_budget_for_litellm_teams(self) -> None: """ Resets the budget for all LiteLLM Internal Teams if their budget has expired """ + await _run_phase_in_chunks(self._reset_budget_for_litellm_teams_chunk) + + async def _reset_budget_for_litellm_teams_chunk(self) -> _ChunkOutcome: now: Final = datetime.utcnow() start_time: Final = time.time() teams_to_reset: list[LiteLLM_TeamTable] | None = None try: - teams_to_reset = await self.prisma_client.get_data(table_name="team", query_type="find_all", reset_at=now) + teams_to_reset = await self.prisma_client.get_data( + table_name="team", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ) updated_teams: Final[list[LiteLLM_TeamTable]] = [] failed_teams: Final = [] if teams_to_reset is not None and len(teams_to_reset) > 0: @@ -688,8 +832,27 @@ class ResetBudgetJob: await self._invalidate_spend_counter(f"spend:team:{team_id}") end_time = time.time() - if len(failed_teams) > 0: # If any teams failed to reset - raise Exception(f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}") + outcome: Final = _ChunkOutcome( + fetched=len(teams_to_reset) if teams_to_reset else 0, + advanced=_count_advanced( + (t.budget_reset_at for t in updated_teams), + cutoff=datetime.now(timezone.utc), + ), + ) + if len(failed_teams) > 0: + self._emit_phase_failure( + call_type="reset_budget_teams", + error=Exception( + f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}" + ), + start_time=start_time, + end_time=end_time, + event_metadata={ + "num_teams_found": len(teams_to_reset) if teams_to_reset else 0, + "teams_found": json.dumps(teams_to_reset, indent=4, default=str), + }, + ) + return outcome asyncio.create_task( self.proxy_logging_obj.service_logging_obj.async_service_success_hook( @@ -725,6 +888,9 @@ class ResetBudgetJob: ) ) verbose_proxy_logger.exception("Failed to reset budget for teams: %s", e) + return _NO_PROGRESS + else: + return outcome @staticmethod async def _reset_expired_window( @@ -882,33 +1048,6 @@ class ResetBudgetJob: ) return user - @staticmethod - async def _reset_budget_for_enduser( - enduser: LiteLLM_EndUserTable, - ) -> LiteLLM_EndUserTable | None: - try: - enduser.spend = 0.0 - except Exception as e: - verbose_proxy_logger.exception("Error resetting budget for enduser: %s. Item: %s", e, enduser) - raise e - return enduser - - @staticmethod - async def _reset_budget_reset_at_date( - budget: LiteLLM_BudgetTableFull, - current_time: datetime, - reset_settings: BudgetResetSettings, - ) -> LiteLLM_BudgetTableFull: - try: - if budget.budget_duration is not None: - budget.budget_reset_at = compute_budget_reset_at( - budget_duration=budget.budget_duration, settings=reset_settings - ) - except Exception as e: - verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget) - raise e - return budget - @staticmethod async def _reset_budget_for_key( key: LiteLLM_VerificationToken, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 53c5112d1d7..3fd2adda480 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -31,6 +31,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( + ESTIMATED_OUTPUT_TOKENS_FIELD, + get_estimated_output_tokens, get_key_tag_rpm_limit, get_model_rate_limit_from_metadata, ) @@ -562,6 +564,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data: dict, model: str | None = None, min_configured_tpm_limit: int | None = None, + configured_output_tokens: int | None = None, ) -> int: """ Estimate total tokens this request will consume so we can reserve them @@ -575,6 +578,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): provided, the no-``max_tokens`` output-budget floor is capped at a fraction of that limit so small TPM caps remain usable. Omit to preserve the unconstrained floor. + + ``configured_output_tokens`` is the operator-declared estimate resolved + from key or team metadata. When provided it replaces the heuristic + floor entirely, so the reservation reflects what this tenant's model + actually emits rather than one constant shared by every tenant. """ messages = data.get("messages") prompt: Final = data.get("prompt") @@ -604,7 +612,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): case (_, embeddings_input) if embeddings_input: # Embeddings have no output tokens max_tokens_estimate = 0 - case _ if total_chars == 0: + case _ if total_chars == 0 and configured_output_tokens is None: # Fully contentless request (no messages, prompt, or input). # Don't apply the conservative output-budget floor here — it # would over-reserve and could push small TPM limits into a @@ -619,7 +627,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # so a small per-tenant TPM cap can't be tripped by the floor # alone. output_floor: Final = self._no_max_tokens_output_floor(min_configured_tpm_limit) - max_tokens_estimate = max(estimated_input_tokens, output_floor) + max_tokens_estimate = ( + configured_output_tokens + if configured_output_tokens is not None + else max(estimated_input_tokens, output_floor) + ) total_estimated: Final = estimated_input_tokens + max_tokens_estimate @@ -2586,8 +2598,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data.get("max_tokens") is not None or data.get("max_completion_tokens") is not None ) is_embedding: Final = data.get("input") is not None + configured_output_tokens: Final = get_estimated_output_tokens( + user_api_key_dict=user_api_key_dict, + model_name=requested_model, + ) if capped_floor < baseline_floor and not has_explicit_max_tokens and not is_embedding: - data["max_tokens"] = capped_floor + data["max_tokens"] = max(capped_floor, configured_output_tokens or 0) # Floor at 1 token so contentless requests (/responses, # tool-call continuations, empty messages) still flow @@ -2601,10 +2617,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data=data, model=requested_model, min_configured_tpm_limit=min_configured_tpm_limit, + configured_output_tokens=configured_output_tokens, ), 1, ) + if configured_output_tokens is not None and estimated_tokens > min_configured_tpm_limit: + verbose_proxy_logger.debug( + "Reserving %s tokens for model %s (declared %s=%s plus the input estimate) exceeds the " + "smallest TPM limit this request is charged against (%s), so it cannot be admitted even " + "against an empty window. Lower the declared estimate or raise the TPM limit.", + estimated_tokens, + requested_model, + ESTIMATED_OUTPUT_TOKENS_FIELD, + configured_output_tokens, + min_configured_tpm_limit, + ) + tpm_response: Final = await self.reserve_tpm_tokens( descriptors=descriptors, estimated_tokens=estimated_tokens, diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 446ea76752e..8c6195388c5 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -20,7 +20,10 @@ from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( + _user_has_admin_view, + validate_budget_duration, +) from litellm.proxy.utils import jsonify_object from litellm.repositories.budget_repository import BudgetRepository @@ -72,6 +75,8 @@ async def new_budget( detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"}, ) + validate_budget_duration(budget_obj.budget_duration) + # Validate model_max_budget if present if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0: from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -153,6 +158,8 @@ async def update_budget( detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"}, ) + validate_budget_duration(budget_obj.budget_duration) + # Validate model_max_budget if present in update if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0: from litellm.proxy.management_endpoints.key_management_endpoints import ( diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index fb302e87bd9..d1542b38996 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -10,6 +10,7 @@ from typing_extensions import TypedDict from litellm._logging import verbose_proxy_logger from litellm.constants import PTU_SENTINEL_API_KEY from litellm.proxy._types import CommonProxyErrors +from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.proxy.utils import PrismaClient from litellm.repositories.table_repositories import DeletedVerificationTokenRepository from litellm.repositories.verification_token_repository import ( @@ -141,6 +142,28 @@ class _GroupingSetsRow(SimpleNamespace): failed_requests: int | None +def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float: + """Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled. + + Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost`` + column straight off the row, and the aggregated path reads the SUM() alias. Rows an + operator accrued during an earlier opt-in stay in the table, so the gate lives on the + read rather than on the query that produced the rows. + + The row is checked before the flag because this runs once per metric accumulation, and + a record fans out across roughly a dozen breakdowns. The flag reads through the secret + manager, uncached, so consulting it for every accumulation put thousands of lookups on + a shared endpoint that made none before. Only a row actually carrying flat cost, which + is a sentinel row, reaches it now. + """ + raw: Final = getattr(record, "ptu_flat_cost", None) or 0.0 + if not raw: + return 0.0 + if not is_ptu_cost_attribution_enabled(): + return 0.0 + return raw + + def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> SpendMetrics: """Update metrics with new record data. @@ -151,7 +174,7 @@ def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> prompt_tokens: Final = record.prompt_tokens or 0 completion_tokens: Final = record.completion_tokens or 0 existing_metrics.spend += record.spend or 0.0 - existing_metrics.flat_cost += getattr(record, "ptu_flat_cost", None) or 0.0 + existing_metrics.flat_cost += _reported_flat_cost(record) existing_metrics.prompt_tokens += prompt_tokens existing_metrics.completion_tokens += completion_tokens existing_metrics.total_tokens += prompt_tokens + completion_tokens @@ -784,7 +807,7 @@ def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics: completion_tokens: Final = record.completion_tokens or 0 return SpendMetrics( spend=record.spend or 0.0, - flat_cost=getattr(record, "ptu_flat_cost", None) or 0.0, + flat_cost=_reported_flat_cost(record), prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens, diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 3868b04f385..2241884faf1 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -22,6 +22,35 @@ def validate_finite_spend(spend: float | None) -> None: ) +def validate_budget_duration(budget_duration: str | None) -> None: + """Reject budget durations that can't be parsed, are non-positive, or + overflow date math, so a bad value can't be persisted and later crash the + budget reset job. + + A non-positive duration also resolves to a reset time of "now", which leaves + the row permanently due: the reset job re-reads it every tick and, once + enough of them exist, they fill each batch and starve every other tenant's + reset. + """ + if budget_duration is None: + return + + from litellm.litellm_core_utils.duration_parser import duration_in_seconds + from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + + try: + if duration_in_seconds(budget_duration) <= 0: + raise ValueError("budget_duration must be positive") + get_budget_reset_time(budget_duration=budget_duration) + except (ValueError, OverflowError): + raise HTTPException( + status_code=400, + detail={ + "error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'." + }, + ) + + from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache from litellm.proxy._types import ( diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index a51ff48aab6..bfc70da46ea 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -23,6 +23,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity +from litellm.proxy.management_endpoints.common_utils import validate_budget_duration from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, handle_update_object_permission_common, @@ -184,6 +185,7 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None: if budget_kv_pairs: budget_request: Final = BudgetNewRequest(**budget_kv_pairs) + validate_budget_duration(budget_request.budget_duration) if budget_request.budget_reset_at is None and budget_request.budget_duration is not None: budget_request.budget_reset_at = datetime.utcnow() + timedelta( seconds=duration_in_seconds(duration=budget_request.budget_duration) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index abc5d3e53ff..a416a197ab8 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -42,6 +42,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _is_user_team_admin, _user_has_admin_view, require_caller_user_id_for_non_admin, + validate_budget_duration, validate_finite_spend, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -506,6 +507,8 @@ async def new_user( status_code=500, detail=CommonProxyErrors.db_not_connected_error.value, ) + validate_budget_duration(data.budget_duration) + # Check for duplicate user_id or email await _check_duplicate_user_id(data.user_id, prisma_client) await _check_duplicate_user_email(data.user_email, prisma_client) @@ -1185,6 +1188,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda if "budget_duration" in non_default_values: from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + validate_budget_duration(non_default_values["budget_duration"]) non_default_values["budget_reset_at"] = get_budget_reset_time( budget_duration=non_default_values["budget_duration"] ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2de1d177b33..38b5d755535 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -55,7 +55,10 @@ from litellm.proxy.auth.auth_checks import ( get_project_object, get_team_object, ) -from litellm.proxy.auth.auth_utils import abbreviate_api_key +from litellm.proxy.auth.auth_utils import ( + abbreviate_api_key, + enforce_output_token_estimates_are_admin_only, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.callback_utils import ( decrypt_callback_vars, @@ -79,6 +82,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _set_object_metadata_field, _team_member_has_permission, _user_has_admin_view, + validate_budget_duration, validate_finite_spend, ) from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -841,12 +845,21 @@ async def _common_key_generation_helper( premium_user=premium_user, ) + validate_budget_duration(data.budget_duration) + if data.throttle_on_budget_exceeded is True and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: raise HTTPException( status_code=403, detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."}, ) + enforce_output_token_estimates_are_admin_only( + data=data, + existing_metadata=None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) + if data.metadata is not None and data.metadata.get("service_account_id") is not None and data.team_id is None: await validate_team_id_used_in_service_account_request( team_id=data.team_id, @@ -1014,7 +1027,7 @@ async def _common_key_generation_helper( # Only set budget_duration on key when explicitly provided. Keys with budget_id # but no explicit budget_duration follow their linked budget tier's schedule; - # reset_budget_for_keys_linked_to_budgets() resets them when the tier resets. + # reset_budget_for_litellm_budget_table() resets them when the tier resets. # This avoids duplicating budget_duration on keys so tier updates apply automatically. if "budget_duration" in data_json: data_json["key_budget_duration"] = data_json.pop("budget_duration", None) @@ -1584,6 +1597,8 @@ async def generate_key_fn( - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. + - default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate. + - default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value. - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. - tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit. - tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". @@ -1793,6 +1808,8 @@ async def generate_service_account_key_fn( - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. + - default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate. + - default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value. - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" @@ -2387,6 +2404,7 @@ async def _validate_update_key_data( """Validate permissions and constraints for key update.""" # Reject NaN/±inf spend before it can reach the DB / spend counter. validate_finite_spend(data.spend) + validate_budget_duration(data.budget_duration) _is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -2473,6 +2491,13 @@ async def _validate_update_key_data( detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."}, ) + enforce_output_token_estimates_are_admin_only( + data=data, + existing_metadata=_existing_metadata if isinstance(_existing_metadata, dict) else None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) + # Personal-key bypass: the caller both created the key AND still owns it # (user_id == caller). Checking only created_by would let a demoted admin # who originally created a key for another user continue editing it without @@ -2655,6 +2680,8 @@ async def update_key_fn( - mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200} - tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit. - model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000} + - default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. + - default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above {"gpt-4": 4096, "gpt-3.5-turbo": 1024} - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - allowed_cache_controls: Optional[list] - List of allowed cache control values @@ -4629,6 +4656,15 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) + if data is not None: + _existing_key_metadata: Final = getattr(key_in_db, "metadata", None) + enforce_output_token_estimates_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, + entity="key", + ) + new_token: Final = await get_new_token(data=data) new_token_hash: Final = hash_token(new_token) new_token_key_name: Final = abbreviate_api_key(api_key=new_token) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index cf6af66d5a9..8a52b0d1abb 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -56,6 +56,10 @@ from litellm.proxy.management_endpoints.team_endpoints import ( update_team as _legacy_update_team, ) from litellm.proxy.management_helpers.audit_logs import create_object_audit_log +from litellm.proxy.spend_tracking.ptu_feature_flag import ( + PTU_COST_ATTRIBUTION_ENV_VAR, + is_ptu_cost_attribution_enabled, +) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ModelTableRepository @@ -239,8 +243,12 @@ _PTU_MODEL_INFO_FIELDS: Final = ("ptu_count", "cost_per_ptu_per_hour", "ptu_effe def _explicitly_cleared_ptu_fields(model_info: ModelInfo | None) -> frozenset[str]: - """The PTU fields a patch sends as an explicit null, which update_db_model drops.""" - if model_info is None: + """The PTU fields a patch sends as an explicit null, which update_db_model drops. + + Empty while the feature is off, so disabling pauses PTU rather than letting a client + that round-trips a model_info blob erase a configuration set up during an earlier opt-in. + """ + if model_info is None or not is_ptu_cost_attribution_enabled(): return frozenset() return frozenset( field @@ -262,6 +270,32 @@ def _merged_ptu_model_info(*, db_model: Deployment, patch_data: updateDeployment return MappingProxyType({k: v for k, v in {**stored, **incoming}.items() if k not in cleared}) +def _raise_if_ptu_cost_attribution_disabled(incoming_model_info: Mapping[str, object]) -> None: + """Reject PTU model_info fields unless the operator opted into PTU cost attribution. + + Takes the incoming request's model_info rather than the merged deployment, so an + unrelated patch of a model that still stores PTU config from an earlier opt-in is + left alone. The fields are rejected rather than dropped so a caller never believes + a flat cost was configured while the rollup that would price it is not running. + + Only a value is rejected. An explicit null reaches the clear loop, which is gated on + the same flag, so a disabled proxy neither writes PTU config nor erases what an + earlier opt-in stored. Disabling pauses the feature rather than discarding its setup. + """ + if is_ptu_cost_attribution_enabled(): + return + supplied: Final = tuple(field for field in _PTU_MODEL_INFO_FIELDS if incoming_model_info.get(field) is not None) + if not supplied: + return + raise HTTPException( + status_code=400, + detail=( + f"PTU cost attribution is disabled, so {', '.join(supplied)} cannot be set. " + f"Set {PTU_COST_ATTRIBUTION_ENV_VAR}=true to enable it." + ), + ) + + def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None: """Enforce the PTU cross-field invariant on the effective model_info. @@ -326,6 +360,8 @@ def _coerce_ptu_datetime(value: object) -> datetime.datetime | None: def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel: + if updated_patch.model_info is not None: + _raise_if_ptu_cost_attribution_disabled(updated_patch.model_info.model_dump(exclude_none=True)) merged_model_name: Final = updated_patch.model_name or db_model.model_name merged_litellm_params: Final = db_model.litellm_params.model_dump(exclude_none=True) merged_model_info: Final = db_model.model_info.model_dump(exclude_none=True) @@ -821,6 +857,7 @@ async def _update_team_model_in_db( # raising the rate on a configured model carries no ptu_effective_from, which the # stored row supplies. if patch_data.model_info is not None: + _raise_if_ptu_cost_attribution_disabled(patch_data.model_info.model_dump(exclude_none=True)) _validate_ptu_model_info(_merged_ptu_model_info(db_model=db_model, patch_data=patch_data)) patch_team_id: Final = patch_data.model_info.team_id if patch_data.model_info else None @@ -1531,7 +1568,9 @@ async def add_new_model( model_response: LiteLLM_ProxyModelTable | None = None # update DB - _validate_ptu_model_info(model_params.model_info.model_dump(exclude_none=True)) + incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True) + _raise_if_ptu_cost_attribution_disabled(incoming_model_info) + _validate_ptu_model_info(incoming_model_info) if store_model_in_db is True: """ diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 97f494c51de..60d3d650d00 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -82,6 +82,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_object, get_user_object, ) +from litellm.proxy.auth.auth_utils import enforce_output_token_estimates_are_admin_only from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch @@ -95,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _update_metadata_fields, _upsert_budget_and_membership, _user_has_admin_view, + validate_budget_duration, ) from litellm.proxy.management_endpoints.organization_endpoints import ( add_member_to_organization, @@ -1153,6 +1155,8 @@ async def new_team( - metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"} - model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team. - model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team. + - default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer. + - default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024} - mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team. - tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit - rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit @@ -1255,6 +1259,9 @@ async def new_team( detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"}, ) + validate_budget_duration(data.budget_duration) + validate_budget_duration(data.team_member_budget_duration) + if data.soft_budget is not None: if data.max_budget is not None: # If max_budget is set, soft_budget must be strictly lower than max_budget @@ -1266,6 +1273,13 @@ async def new_team( }, ) + enforce_output_token_estimates_are_admin_only( + data=data, + existing_metadata=None, + user_api_key_dict=user_api_key_dict, + entity="team", + ) + # Check if license is over limit total_teams: Final = await _team_db(prisma_client).count() if total_teams and _license_check.is_team_count_over_limit(team_count=total_teams): @@ -1863,6 +1877,8 @@ async def update_team( - allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team. - model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit per model for this team. Example: {"gpt-4": 100, "gpt-3.5-turbo": 200} - model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit per model for this team. Example: {"gpt-4": 10000, "gpt-3.5-turbo": 20000} + - default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer. + - default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024} - mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team. Example - update team TPM Limit - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. @@ -1935,6 +1951,9 @@ async def update_team( detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"}, ) + validate_budget_duration(data.budget_duration) + validate_budget_duration(data.team_member_budget_duration) + existing_team_row = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id}) if existing_team_row is None: @@ -1949,6 +1968,14 @@ async def update_team( user_api_key_dict=user_api_key_dict, ) + _existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None) + enforce_output_token_estimates_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, + entity="team", + ) + _check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") if data.soft_budget is not None: @@ -2959,7 +2986,7 @@ async def team_member_add( except HTTPException as e: raise e - _validate_budget_duration(data.budget_duration) + validate_budget_duration(data.budget_duration) prisma_client = cast(PrismaClient, prisma_client) @@ -3262,29 +3289,6 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, objec } -def _validate_budget_duration(budget_duration: str | None) -> None: - """Reject budget durations that can't be parsed, are non-positive, or - overflow date math, so a bad value can't be persisted and later crash the - budget reset job.""" - if budget_duration is None: - return - - from litellm.litellm_core_utils.duration_parser import duration_in_seconds - from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time - - try: - if duration_in_seconds(budget_duration) <= 0: - raise ValueError("budget_duration must be positive") - get_budget_reset_time(budget_duration=budget_duration) - except (ValueError, OverflowError): - raise HTTPException( - status_code=400, - detail={ - "error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'." - }, - ) - - @router.post( "/team/member_update", tags=["team management"], @@ -3322,7 +3326,7 @@ async def team_member_update( detail={"error": "Either user_id or user_email needs to be passed in"}, ) - _validate_budget_duration(data.budget_duration) + validate_budget_duration(data.budget_duration) _existing_team_row: Final = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id}) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 641d294f592..bc980934f9f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6860,9 +6860,19 @@ class ProxyConfig: guardrail_id = guardrail.get("guardrail_id") if guardrail_id: db_guardrail_ids.add(guardrail_id) - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( - guardrail=cast(Guardrail, guardrail), - ) + try: + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( + guardrail=cast(Guardrail, guardrail), + ) + except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails + verbose_proxy_logger.error( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - " + "skipping guardrail '%s' (ID: %s): %s: %s", + guardrail.get("guardrail_name"), + guardrail_id, + type(e).__name__, + e, + ) # Drop in-memory DB-backed entries whose row was deleted on another # pod. Config-loaded entries are never touched. @@ -8471,40 +8481,45 @@ class ProxyStartupEvent: await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler) ### PTU DAILY ROLLUP ### - from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( - PTU_ROLLUP_JOB_ID, - run_scheduled_ptu_rollup, + from litellm.proxy.spend_tracking.ptu_feature_flag import ( + is_ptu_cost_attribution_enabled, ) - async def _alert_ptu_rollup_failure(message: str) -> None: - await proxy_logging_obj.alerting_handler( - message=message, - level="High", - alert_type=AlertType.failed_tracking_spend, + if is_ptu_cost_attribution_enabled(): + from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( + PTU_ROLLUP_JOB_ID, + run_scheduled_ptu_rollup, ) - async def _scheduled_ptu_rollup() -> None: - # Reuse the PodLockManager from db_spend_update_writer so only one pod - # reconciles a day; a multi-pod race could prune another pod's fresh rows - await run_scheduled_ptu_rollup( - prisma_client, - pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager, - alert=_alert_ptu_rollup_failure, - ) + async def _alert_ptu_rollup_failure(message: str) -> None: + await proxy_logging_obj.alerting_handler( + message=message, + level="High", + alert_type=AlertType.failed_tracking_spend, + ) - scheduler.add_job( - _scheduled_ptu_rollup, - "cron", - hour=0, - minute=15, - timezone="UTC", - id=PTU_ROLLUP_JOB_ID, - replace_existing=True, - misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, - ) - verbose_proxy_logger.info( - "PTU rollup job scheduled at 00:15 UTC daily (only models with PTU config accrue flat cost)" - ) + async def _scheduled_ptu_rollup() -> None: + # Reuse the PodLockManager from db_spend_update_writer so only one pod + # reconciles a day; a multi-pod race could prune another pod's fresh rows + await run_scheduled_ptu_rollup( + prisma_client, + pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager, + alert=_alert_ptu_rollup_failure, + ) + + scheduler.add_job( + _scheduled_ptu_rollup, + "cron", + hour=0, + minute=15, + timezone="UTC", + id=PTU_ROLLUP_JOB_ID, + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) + verbose_proxy_logger.info( + "PTU rollup job scheduled at 00:15 UTC daily (only models with PTU config accrue flat cost)" + ) ### SPEND LOG CLEANUP ### if ( diff --git a/litellm/proxy/spend_tracking/ptu_feature_flag.py b/litellm/proxy/spend_tracking/ptu_feature_flag.py new file mode 100644 index 00000000000..9078079b676 --- /dev/null +++ b/litellm/proxy/spend_tracking/ptu_feature_flag.py @@ -0,0 +1,18 @@ +"""Opt-in flag for PTU (provisioned throughput unit) flat-cost attribution. + +The whole feature is inert unless an operator sets +``LITELLM_ENABLE_PTU_COST_ATTRIBUTION``: the daily rollup is not scheduled, the +model endpoints reject PTU config, the daily activity read path reports zero flat +cost, and the model form hides the PTU inputs. +""" + +from typing import Final + +from litellm.secret_managers.main import get_secret_bool + +PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION" + + +def is_ptu_cost_attribution_enabled() -> bool: + """Report whether this deployment opted into PTU flat-cost attribution.""" + return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index a77c3e29cfa..029648f7901 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -27,6 +27,7 @@ from litellm.constants import ( PTU_ROLLUP_MAX_BACKFILL_DAYS, PTU_SENTINEL_API_KEY, ) +from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.types.router import ModelInfo if TYPE_CHECKING: @@ -512,7 +513,15 @@ async def run_scheduled_ptu_rollup( duplicate work rather than correctness: the upserts are idempotent on the sentinel key and the prune reads only the row's own timestamp, so a second pod arriving mid-run cannot corrupt the day. + + Returns None without touching the database when PTU cost attribution is off. Proxy + startup already skips scheduling the cron, so this guards the function itself rather + than its one caller, and a deployment that never opted in accrues nothing whatever + reaches it. """ + if not is_ptu_cost_attribution_enabled(): + return None + if pod_lock_manager is None or pod_lock_manager.redis_cache is None: return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 08bb8698cac..b3feb5bd8d6 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -21,6 +21,7 @@ from litellm.proxy.config_resolvers.sso import ( SSO_SECRET_FIELDS, resolve_sso_config, ) +from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.proxy.utils import invalidate_config_param from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.organization_repository import OrganizationRepository @@ -307,6 +308,27 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = { "enable_chat_ui", } +ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: Final = "enable_ptu_cost_attribution" + +# UI settings derived from the deployment environment. Deliberately kept out of +# ALLOWED_UI_SETTINGS_FIELDS: they are read-only, never persisted, and PATCH +# rejects them so an admin cannot flip an env-gated feature at runtime. +_DERIVED_UI_SETTINGS_FIELDS: Final[frozenset[str]] = frozenset({ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING}) + + +def _derived_ui_setting_value(key: str) -> object: + """The environment-derived value GET reports for ``key``. + + PATCH compares against this rather than rejecting the key outright, so the body GET + hands back is still a valid PATCH body. Rejecting on presence broke read-modify-write: + a client that edited one setting and sent the rest back unchanged got a 400 and lost + the edit it actually wanted. + """ + if key == ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: + return is_ptu_cost_attribution_enabled() + return None + + # Flags that must be synced from the persisted UISettings into # general_settings at runtime (on both read and write). _RUNTIME_GENERAL_SETTINGS_FLAGS: Final = [ @@ -1345,21 +1367,15 @@ async def get_ui_settings(): detail={"error": "Database not connected. Please connect a database."}, ) - ui_settings: Mapping[str, JsonValue] = {} - db_record: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique( where={"id": "ui_settings"} ) - if db_record and db_record.ui_settings: - ui_settings_json: Final = db_record.ui_settings - if isinstance(ui_settings_json, str): - ui_settings = json.loads(ui_settings_json) - else: - ui_settings = dict(ui_settings_json) + stored: Final = (db_record.ui_settings if db_record else None) or "{}" + parsed: Final = json.loads(stored) if isinstance(stored, str) else stored # Sanitize any unexpected keys from persisted config before returning - ui_settings = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS} + ui_settings: Final = {k: v for k, v in parsed.items() if k in ALLOWED_UI_SETTINGS_FIELDS} # Sync runtime flags into general_settings so the proxy picks them up # at runtime (covers server restart scenarios). @@ -1377,11 +1393,18 @@ async def get_ui_settings(): # Build config-like object for schema helper config: Final[dict[str, object]] = {"litellm_settings": {"ui_settings": ui_settings}} - return await _get_settings_with_schema( + settings: Final = await _get_settings_with_schema( settings_key="ui_settings", settings_class=_get_effective_ui_settings_class(), config=config, ) + return UISettingsResponse( + values={ + **settings["values"], + ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: is_ptu_cost_attribution_enabled(), + }, + field_schema=settings["field_schema"], + ) @router.patch( @@ -1418,6 +1441,20 @@ async def update_ui_settings( detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, ) + conflicting_keys: Final = sorted( + key + for key, value in settings_body.items() + if key in _DERIVED_UI_SETTINGS_FIELDS and value != _derived_ui_setting_value(key) + ) + if conflicting_keys: + raise HTTPException( + status_code=400, + detail=( + f"Setting(s) {conflicting_keys} are derived from the deployment environment " + "and cannot be changed from the UI." + ), + ) + # Validate against the same effective class GET advertises, so # enterprise-registered fields are typed consistently on both sides. effective_cls: Final = _get_effective_ui_settings_class() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5f22ca021ac..dd0c57aa911 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3486,13 +3486,15 @@ class PrismaClient: r.expires = r.expires.isoformat() elif query_type == "find_all" and expires is not None and reset_at is not None: response = await VerificationTokenRepository(self).table.find_many( + take=limit, where={ "OR": [ {"expires": None}, {"expires": {"gt": expires}}, ], "budget_reset_at": {"lt": reset_at}, - } + "NOT": {"budget_duration": None}, + }, ) if response is not None and len(response) > 0: for r in response: @@ -3542,6 +3544,7 @@ class PrismaClient: response = await UserRepository(self).table.find_many(where=key_val) elif query_type == "find_all" and reset_at is not None: response = await UserRepository(self).table.find_many( + take=limit, where={ # A user seeded from default_internal_user_params # (or created via /user/new without an explicit @@ -3552,16 +3555,12 @@ class PrismaClient: # of the row, silently exceeding max_budget. Treat a # NULL budget_reset_at with a non-NULL budget_duration # as due, matching the budget-table query below. + "NOT": {"budget_duration": None}, "OR": [ - { - "AND": [ - {"budget_reset_at": None}, - {"NOT": {"budget_duration": None}}, - ] - }, + {"budget_reset_at": None}, {"budget_reset_at": {"lt": reset_at}}, ], - } + }, ) elif query_type == "find_all" and user_id_list is not None: response = await UserRepository(self).table.find_many(where={"user_id": {"in": user_id_list}}) @@ -3617,17 +3616,14 @@ class PrismaClient: elif table_name == "budget" and reset_at is not None: if query_type == "find_all": response = await BudgetRepository(self).table.find_many( + take=limit, where={ + "NOT": {"budget_duration": None}, "OR": [ - { - "AND": [ - {"budget_reset_at": None}, - {"NOT": {"budget_duration": None}}, - ] - }, + {"budget_reset_at": None}, {"budget_reset_at": {"lt": reset_at}}, - ] - } + ], + }, ) return response @@ -3645,20 +3641,17 @@ class PrismaClient: ) elif query_type == "find_all" and reset_at is not None: response = await TeamRepository(self).table.find_many( + take=limit, where={ # Same NULL budget_reset_at gap as the user query # above: a team with a budget_duration but no # initialized budget_reset_at would never be reset. + "NOT": {"budget_duration": None}, "OR": [ - { - "AND": [ - {"budget_reset_at": None}, - {"NOT": {"budget_duration": None}}, - ] - }, + {"budget_reset_at": None}, {"budget_reset_at": {"lt": reset_at}}, ], - } + }, ) elif query_type == "find_all" and user_id is not None: response = await TeamRepository(self).table.find_many( diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 4f020480f9e..e2e7f1fac73 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -70,10 +70,14 @@ from litellm.repositories.table_repositories import ( ) from litellm.repositories.team_repository import TeamRepository from litellm.repositories.unit_of_work import ( + BudgetCascadeUnitOfWork, + BudgetWindowWrites, KeySpendResetWrites, + LinkedSpendResetWrites, SpendResetUnitOfWork, TeamSpendResetWrites, UserSpendResetWrites, + budget_cascade_unit_of_work, spend_reset_unit_of_work, ) from litellm.repositories.user_repository import UserRepository @@ -88,7 +92,9 @@ __all__ = [ "AgentsRepository", "AuditLogRepository", "BatchTable", + "BudgetCascadeUnitOfWork", "BudgetRepository", + "BudgetWindowWrites", "CacheConfigRepository", "ClaudeCodePluginRepository", "ConfigOverridesRepository", @@ -107,6 +113,7 @@ __all__ = [ "InvitationLinkRepository", "JWTKeyMappingRepository", "KeySpendResetWrites", + "LinkedSpendResetWrites", "MCPServerRepository", "MCPToolsetRepository", "MCPUserCredentialsRepository", @@ -149,5 +156,6 @@ __all__ = [ "WorkflowEventRepository", "WorkflowMessageRepository", "WorkflowRunRepository", + "budget_cascade_unit_of_work", "spend_reset_unit_of_work", ] diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 6aff196ff10..055c68163f9 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -29,6 +29,8 @@ class SpendLinkedTable(Protocol[RowT_co]): class BatchTable(Protocol): def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + class PrismaBatch(Protocol): @property @@ -40,4 +42,19 @@ class PrismaBatch(Protocol): @property def litellm_teamtable(self) -> BatchTable: ... + @property + def litellm_budgettable(self) -> BatchTable: ... + + @property + def litellm_teammembership(self) -> BatchTable: ... + + @property + def litellm_organizationtable(self) -> BatchTable: ... + + @property + def litellm_tagtable(self) -> BatchTable: ... + + @property + def litellm_endusertable(self) -> BatchTable: ... + async def commit(self) -> None: ... diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py index 682e69d11eb..e504baceb9f 100644 --- a/litellm/repositories/unit_of_work.py +++ b/litellm/repositories/unit_of_work.py @@ -1,17 +1,21 @@ """ -Unit of work over a single Prisma batch. +Units of work over a single Prisma batch. -``spend_reset_unit_of_work`` opens one ``db.batch_()`` and binds a typed write +Each context manager here opens one ``db.batch_()`` and binds a typed write repository per table to it, so every update queued through the yielded object lands in the same transaction. The batch commits when the block exits cleanly and is abandoned, writing nothing, when the block raises. -Each write repository queues narrow ``{spend, budget_reset_at}`` updates +``spend_reset_unit_of_work`` covers the per-row key/user/team resets; +``budget_cascade_unit_of_work`` covers a budget tier's reset, where the +dependent spend and the tier's next window have to move together. + +Each write repository queues narrow ``{spend}`` / ``{budget_reset_at}`` updates instead of full-model writes, which trip ``prisma.errors.DataError`` on rows carrying fields the update input type rejects (see #27730). """ -from collections.abc import AsyncGenerator, Callable +from collections.abc import AsyncGenerator, Callable, Mapping from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import datetime @@ -43,6 +47,24 @@ class TeamSpendResetWrites: self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at}) +@dataclass(frozen=True, slots=True) +class LinkedSpendResetWrites: + table: BatchTable + + def queue_spend_zero(self, where: Mapping[str, object]) -> None: + self.table.update_many(where=where, data={"spend": 0}) + + +@dataclass(frozen=True, slots=True) +class BudgetWindowWrites: + table: BatchTable + + def queue_window_advance(self, budget_id: str, budget_reset_at: datetime) -> None: + """``update_many`` so a tier deleted between the read and the commit is a + no-op row count instead of a P2025 that aborts the whole chunk.""" + self.table.update_many(where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at}) + + @dataclass(frozen=True, slots=True) class SpendResetUnitOfWork: keys: KeySpendResetWrites @@ -50,6 +72,23 @@ class SpendResetUnitOfWork: teams: TeamSpendResetWrites +@dataclass(frozen=True, slots=True) +class BudgetCascadeUnitOfWork: + """Every write a budget-tier reset performs, bound to one batch. + + The dependent spend rows and the budget rows' ``budget_reset_at`` advance + must land together: advancing the window without zeroing the spend it + gates leaves the dependents pinned at their cap until the next window. + """ + + team_memberships: LinkedSpendResetWrites + keys: LinkedSpendResetWrites + organizations: LinkedSpendResetWrites + tags: LinkedSpendResetWrites + endusers: LinkedSpendResetWrites + budgets: BudgetWindowWrites + + @asynccontextmanager async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> AsyncGenerator[SpendResetUnitOfWork, None]: batch = new_batch() @@ -59,3 +98,19 @@ async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> Asyn teams=TeamSpendResetWrites(table=batch.litellm_teamtable), ) await batch.commit() + + +@asynccontextmanager +async def budget_cascade_unit_of_work( + new_batch: Callable[[], PrismaBatch], +) -> AsyncGenerator[BudgetCascadeUnitOfWork, None]: + batch = new_batch() + yield BudgetCascadeUnitOfWork( + team_memberships=LinkedSpendResetWrites(table=batch.litellm_teammembership), + keys=LinkedSpendResetWrites(table=batch.litellm_verificationtoken), + organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable), + tags=LinkedSpendResetWrites(table=batch.litellm_tagtable), + endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable), + budgets=BudgetWindowWrites(table=batch.litellm_budgettable), + ) + await batch.commit() diff --git a/osv-scanner.toml b/osv-scanner.toml index 4ef612e3a70..7ab450945f5 100644 --- a/osv-scanner.toml +++ b/osv-scanner.toml @@ -1,13 +1,3 @@ -[[IgnoredVulns]] -id = "GHSA-fwg2-594c-jp42" -ignoreUntil = 2026-08-12 -reason = "pypdf 6.15.0 (the fix) published 2026-08-06 and is still inside the P3D exclude-newer window, so uv cannot lock it yet; bump and drop this entry from 2026-08-09" - -[[IgnoredVulns]] -id = "GHSA-fp3f-mc75-235c" -ignoreUntil = 2026-08-12 -reason = "second pypdf advisory with the same 6.15.0 fix, published 2026-08-07 after the first; drop alongside GHSA-fwg2-594c-jp42 in the same bump" - [[IgnoredVulns]] id = "GHSA-w8v5-vhqr-4h9v" ignoreUntil = 2026-09-09 diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 4d1e73aab2d..da0f608fdb5 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -9,10 +9,10 @@ "limit": 832 }, "ANN201": { - "limit": 2031 + "limit": 2023 }, "ANN202": { - "limit": 861 + "limit": 860 }, "ANN204": { "limit": 713 @@ -237,13 +237,13 @@ "limit": 1226 }, "TRY002": { - "limit": 528 + "limit": 524 }, "TRY004": { "limit": 96 }, "TRY201": { - "limit": 407 + "limit": 405 }, "TRY203": { "limit": 113 diff --git a/tests/code_coverage_tests/check_prisma_binary_cache.py b/tests/code_coverage_tests/check_prisma_binary_cache.py new file mode 100644 index 00000000000..501688385ff --- /dev/null +++ b/tests/code_coverage_tests/check_prisma_binary_cache.py @@ -0,0 +1,143 @@ +"""Guard the CI cache for Prisma's CLI and engine binaries. + +``prisma generate`` shells out to ``npm install prisma@`` whenever the +prisma-client-py binary cache directory has no CLI entrypoint, pulling ~85 MB of +engines over the network. The download is normally seconds and occasionally +minutes, and a job timeout cannot tell the difference from a hung test, so an +uncached job is one slow npm response away from cancelling a passing test run. + +Three invariants keep that download off the critical path: + +1. No workflow sets ``PRISMA_BINARY_CACHE_DIR``. The prisma-client-py default is + ``~/.cache/prisma-python/binaries//``, already + keyed by both versions and the only path the cache action restores. Pointing + it elsewhere (``runner.temp`` especially, which is wiped every job) silently + guarantees a cold download. +2. Every job that generates the client also restores the cache. +3. The cache key resolves to a real version from ``uv.lock``. The action fails + the job when it cannot, so a lock format change must break here instead. +""" + +import re +import sys +from collections.abc import Iterator, Mapping +from pathlib import Path +from typing import Final + +import yaml +from pydantic import BaseModel, Field, ValidationError + +REPO_ROOT: Final = Path(__file__).resolve().parent.parent.parent +WORKFLOWS_DIR: Final = REPO_ROOT / ".github" / "workflows" +UV_LOCK: Final = REPO_ROOT / "uv.lock" +CACHE_ACTION: Final = "./.github/actions/cache-prisma-binaries" + +# Commands that reach the prisma binary cache: a direct generate, or a script +# that runs one on the caller's behalf. +PRISMA_GENERATE_MARKERS: Final = ("prisma generate", "type_check_gate.py") + + +class PrismaBinaryCacheError(Exception): + pass + + +def resolve_prisma_version(lock_text: str) -> str | None: + """Mirror of the shell lookup in the cache action's version step.""" + match: Final = re.search( + r'^name = "prisma"\n^version = "(?P[^"]+)"$', + lock_text, + re.MULTILINE, + ) + return match.group("version") if match else None + + +class WorkflowStep(BaseModel): + """The two step fields this guard reads; every other key is ignored.""" + + run: str | None = None + uses: str | None = None + + def generates_prisma_client(self) -> bool: + return self.run is not None and any(m in self.run for m in PRISMA_GENERATE_MARKERS) + + def restores_cache(self) -> bool: + return self.uses == CACHE_ACTION + + +class WorkflowJob(BaseModel): + # Absent for jobs that delegate to a reusable workflow via a job-level `uses`. + steps: tuple[WorkflowStep, ...] = () + + +class Workflow(BaseModel): + jobs: Mapping[str, WorkflowJob] = Field(default_factory=dict) + + +def parse_workflow(text: str) -> Workflow | str: + """Validate untyped YAML at the boundary so the checks below stay typed. + + Returns the parsed workflow, or a description of why it could not be read. + """ + parsed: Final = yaml.safe_load(text) + try: + return Workflow.model_validate(parsed if isinstance(parsed, dict) else {}) + except ValidationError as exc: + return f"does not parse as a workflow: {exc.error_count()} schema error(s)" + + +def lock_errors(lock_text: str) -> Iterator[str]: + if not resolve_prisma_version(lock_text): + yield ( + "uv.lock has no resolvable `prisma` package version. The version step " + f"in {CACHE_ACTION} greps the same shape and will fail every job that " + "generates the Prisma client." + ) + + +def workflow_errors(rel: Path, text: str) -> Iterator[str]: + if "PRISMA_BINARY_CACHE_DIR" in text: + yield ( + f"{rel}: sets PRISMA_BINARY_CACHE_DIR. Leave it unset so the binaries " + f"land in the version-keyed default path the {CACHE_ACTION} action restores." + ) + + workflow: Final = parse_workflow(text) + if isinstance(workflow, str): + yield f"{rel}: {workflow}" + return + + for job_name, job in workflow.jobs.items(): + if any(s.generates_prisma_client() for s in job.steps) and not any( + s.restores_cache() for s in job.steps + ): + yield ( + f"{rel}: job `{job_name}` generates the Prisma client without a " + f"`uses: {CACHE_ACTION}` step, so it downloads ~85 MB of engines " + "on every run." + ) + + +def main() -> None: + errors: Final = ( + *lock_errors(UV_LOCK.read_text()), + *( + error + for path in sorted(WORKFLOWS_DIR.glob("*.y*ml")) + for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text()) + ), + ) + + if errors: + raise PrismaBinaryCacheError( + "Prisma binary cache invariants violated:\n - " + "\n - ".join(errors) + ) + + print("Prisma binary cache invariants hold across .github/workflows/") + + +if __name__ == "__main__": + try: + main() + except PrismaBinaryCacheError as exc: + print(f"ERROR: {exc}", file=sys.stderr) + sys.exit(1) diff --git a/tests/code_coverage_tests/check_workflow_startup_safety.py b/tests/code_coverage_tests/check_workflow_startup_safety.py new file mode 100644 index 00000000000..cf150daef4c --- /dev/null +++ b/tests/code_coverage_tests/check_workflow_startup_safety.py @@ -0,0 +1,239 @@ +"""Catch workflow mistakes that GitHub reports as nothing at all. + +A workflow whose YAML is valid but whose expressions are not fails at *startup*: +the run is marked failed, no jobs are created, and no check run is ever posted. +Nothing turns red on the PR, so an entire test suite can silently stop running +while the checks list stays green. These invariants have to be enforced here +because CI cannot enforce them on itself. + +1. No arithmetic inside ``${{ }}``. GitHub expressions support grouping, index, + dereference, ``!``, the comparisons, ``&&`` and ``||``, and nothing else. A + ``${{ a + b }}`` is a startup failure, not a value. Only ``+`` and ``*`` are + flagged: ``-`` appears in hyphenated input names like ``inputs.timeout-minutes`` + and ``/`` inside ref strings, so neither can be told apart from arithmetic by + inspection alone. +2. Callers of the reusable unit-test workflow keep the job timeout at or above + the test budget plus the setup ceilings plus the runner overhead below. + Otherwise the job deadline preempts pytest inside its own advertised budget, + which is the failure the split timeouts exist to prevent, and it shows up as + a cancelled shard whose tests were passing. A budget this check cannot resolve + is reported rather than skipped, so a mistyped input or matrix column surfaces + here instead of leaving the pair silently unchecked. +""" + +import re +import sys +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import yaml +from pydantic import BaseModel, Field, ValidationError + +REPO_ROOT: Final = Path(__file__).resolve().parent.parent.parent +WORKFLOWS_DIR: Final = REPO_ROOT / ".github" / "workflows" +BASE_WORKFLOW: Final = "./.github/workflows/_test-unit-base.yml" +BASE_WORKFLOW_PATH: Final = WORKFLOWS_DIR / "_test-unit-base.yml" + +# Runner time the job clock charges but no step owns: job init, the gaps between +# steps, and post-job cleanup. Without it a job capped at exactly test + setup +# would still preempt pytest inside its own budget. +JOB_OVERHEAD_MINUTES: Final = 5 + +EXPRESSION: Final = re.compile(r"\$\{\{(?P.*?)\}\}", re.DOTALL) +QUOTED: Final = re.compile(r"'[^']*'") +ARITHMETIC: Final = re.compile(r"[+*]") +MATRIX_REF: Final = re.compile(r"^\$\{\{\s*matrix\.(?P[\w-]+)\s*\}\}$") + + +class WorkflowStartupError(Exception): + pass + + +class ReusableCall(BaseModel): + uses: str | None = None + with_: Mapping[str, object] = Field(default_factory=dict, alias="with") + strategy: Mapping[str, object] = Field(default_factory=dict) + steps: tuple[Mapping[str, object], ...] = () + + model_config = {"populate_by_name": True} + + +class WorkflowFile(BaseModel): + jobs: Mapping[str, ReusableCall] = Field(default_factory=dict) + + +def parse_workflow(text: str) -> WorkflowFile | str: + parsed: Final = yaml.safe_load(text) + try: + return WorkflowFile.model_validate(parsed if isinstance(parsed, dict) else {}) + except ValidationError as exc: + return f"does not parse as a workflow: {exc.error_count()} schema error(s)" + + +def arithmetic_expressions(text: str) -> Iterator[str]: + for match in EXPRESSION.finditer(text): + body: Final = match.group("body") + if ARITHMETIC.search(QUOTED.sub("", body)): + yield body.strip() + + +def setup_ceiling_minutes(base_text: str) -> int: + """Sum the per-step timeouts on everything the base workflow runs before pytest.""" + base: Final = yaml.safe_load(base_text) + steps: Final = base["jobs"]["run"]["steps"] + return sum( + s["timeout-minutes"] + for s in steps + if s.get("name") != "Run tests" and isinstance(s.get("timeout-minutes"), int) + ) + + +def base_default(base_text: str, name: str) -> int: + base: Final = yaml.safe_load(base_text) + return base[True]["workflow_call"]["inputs"][name]["default"] + + +@dataclass(frozen=True, slots=True) +class Column: + """A budget the caller reads from one column of its own matrix.""" + + name: str + + +def budget_source(job: ReusableCall, key: str, fallback: int) -> int | Column | str: + """A caller passes a literal, or `${{ matrix.x }}` naming a column of its matrix. + + Anything else comes back as the reason it could not be read, since a budget + nothing can resolve has to be reported rather than passed over. + """ + value: Final = job.with_.get(key) + if value is None: + return fallback + if isinstance(value, int): + return value + + matrix_ref: Final = MATRIX_REF.match(str(value)) + if not matrix_ref: + return f"passes `{key}: {value}`, which is neither a number nor a `matrix` reference." + return Column(matrix_ref.group("key")) + + +def matrix_rows(job: ReusableCall) -> Sequence[Mapping[str, object]]: + matrix: Final = job.strategy.get("matrix", {}) + entries: Final = matrix.get("include", ()) if isinstance(matrix, dict) else () + return tuple(e for e in entries if isinstance(e, dict)) + + +def budget_pairs(job: ReusableCall, test_source: int | Column, job_source: int | Column) -> Iterator[tuple[int, int]]: + """Pair each shard's test budget with the job budget of that same shard. + + Matrix-sourced budgets resolve per `include` row, so two matrix columns are + read off the same row rather than cross-producted across rows. + """ + if isinstance(test_source, int) and isinstance(job_source, int): + yield test_source, job_source + return + + for row in matrix_rows(job): + test_budget = row.get(test_source.name) if isinstance(test_source, Column) else test_source + job_budget = row.get(job_source.name) if isinstance(job_source, Column) else job_source + if isinstance(test_budget, int) and isinstance(job_budget, int): + yield test_budget, job_budget + + +def unresolved_message(where: str, job: ReusableCall, sources: Sequence[int | Column]) -> str: + """Why no shard yielded a pair of budgets to compare. + + Naming only the columns that resolve nowhere keeps the message honest: a + column every row supplies is not what left the pair unchecked. + """ + rows: Final = matrix_rows(job) + missing: Final = tuple( + f"`matrix.{s.name}`" + for s in sources + if isinstance(s, Column) and not any(isinstance(row.get(s.name), int) for row in rows) + ) + if missing: + return ( + f"{where} reads a budget from {', '.join(missing)}, which no `include` row supplies " + "as a number, so the pair would go unchecked." + ) + return ( + f"{where} reads both budgets from its matrix, but no single `include` row supplies both " + "as numbers, so the pair would go unchecked." + ) + + +def job_errors(rel: Path, job_name: str, job: ReusableCall, ceiling: int, base_text: str) -> Iterator[str]: + where: Final = f"{rel}: job `{job_name}`" + test_source: Final = budget_source(job, "timeout-minutes", base_default(base_text, "timeout-minutes")) + job_source: Final = budget_source(job, "job-timeout-minutes", base_default(base_text, "job-timeout-minutes")) + sources: Final = (test_source, job_source) + + unreadable: Final = tuple(f"{where} {reason}" for reason in sources if isinstance(reason, str)) + if unreadable: + yield from unreadable + return + + pairs: Final = tuple(budget_pairs(job, test_source, job_source)) + if not pairs: + yield unresolved_message(where, job, sources) + return + + for test_budget, job_budget in pairs: + required = test_budget + ceiling + JOB_OVERHEAD_MINUTES + if job_budget < required: + yield ( + f"{where} gives pytest {test_budget}m but caps the job at " + f"{job_budget}m. Setup can use up to {ceiling}m plus {JOB_OVERHEAD_MINUTES}m of " + f"runner overhead, so the job deadline would preempt pytest; raise " + f"job-timeout-minutes to at least {required}." + ) + + +def timeout_contract_errors(rel: Path, workflow: WorkflowFile, ceiling: int, base_text: str) -> Iterator[str]: + for job_name, job in workflow.jobs.items(): + if job.uses == BASE_WORKFLOW: + yield from job_errors(rel, job_name, job, ceiling, base_text) + + +def workflow_errors(rel: Path, text: str, ceiling: int, base_text: str) -> Iterator[str]: + for expression in arithmetic_expressions(text): + yield ( + f"{rel}: `${{{{ {expression} }}}}` uses arithmetic, which GitHub expressions do not " + "support. The workflow will fail at startup with no jobs and no check run." + ) + + workflow: Final = parse_workflow(text) + if isinstance(workflow, str): + yield f"{rel}: {workflow}" + return + + yield from timeout_contract_errors(rel, workflow, ceiling, base_text) + + +def main() -> None: + base_text: Final = BASE_WORKFLOW_PATH.read_text() + ceiling: Final = setup_ceiling_minutes(base_text) + errors: Final = tuple( + error + for path in sorted(WORKFLOWS_DIR.glob("*.y*ml")) + for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text(), ceiling, base_text) + ) + + if errors: + raise WorkflowStartupError( + "Workflow startup invariants violated:\n - " + "\n - ".join(errors) + ) + + print(f"Workflow startup invariants hold (setup ceiling {ceiling}m)") + + +if __name__ == "__main__": + try: + main() + except WorkflowStartupError as exc: + print(f"ERROR: {exc}", file=sys.stderr) + sys.exit(1) diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index 00d5380b2f4..b13b7342c25 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -44,39 +44,87 @@ def _attrify(d: dict): return _AttrDict(d) -def _wire_batcher_for_test(prisma_client): +def _wire_batcher_for_test(prisma_client, fail_commit=False): """ Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is - awaitable and whose per-table .update() calls get captured. The reset job - writes key/user/team resets via prisma.db.batch_()..update — not via - prisma_client.update_data — so tests must let that batch path complete. + awaitable and whose per-table .update()/.update_many() calls get captured. + The reset job writes every reset through prisma.db.batch_() — key/user/team + rows one by one, and the budget tier's cascade as a single transaction — so + tests must let that batch path complete. - Returns the list that will accumulate {table, where, data} dicts from - each captured update call. + Only committed batches contribute to the returned list, mirroring prisma: + with fail_commit=True the transaction blows up and must persist nothing. + + Returns the list that will accumulate {table, op, where, data} dicts from + each captured write. """ batch_calls = [] def make_batcher(): + queued = [] + class _Table: def __init__(self, table_name): self._table_name = table_name def update(self, where=None, data=None): - batch_calls.append( - {"table": self._table_name, "where": where, "data": data} + queued.append( + { + "table": self._table_name, + "op": "update", + "where": where, + "data": data, + } ) + def update_many(self, where=None, data=None): + queued.append( + { + "table": self._table_name, + "op": "update_many", + "where": where, + "data": data, + } + ) + + async def commit(): + if fail_commit: + raise RuntimeError("simulated Postgres failure committing the batch") + batch_calls.extend(queued) + batcher = MagicMock() batcher.litellm_verificationtoken = _Table("key") batcher.litellm_usertable = _Table("user") batcher.litellm_teamtable = _Table("team") - batcher.commit = AsyncMock(return_value=None) + batcher.litellm_budgettable = _Table("budget") + batcher.litellm_teammembership = _Table("team_membership") + batcher.litellm_organizationtable = _Table("org") + batcher.litellm_tagtable = _Table("tag") + batcher.litellm_endusertable = _Table("enduser") + batcher.commit = commit return batcher prisma_client.db.batch_ = MagicMock(side_effect=make_batcher) return batch_calls +def _wire_cascade_reads_for_test(prisma_client): + """ + The budget tier's cascade reads the rows it is about to zero, so their + spend counters can be invalidated after the commit. Give each of those + tables an awaitable find_many so the reads resolve instead of falling into + the job's warn-and-continue path. + """ + for table in ( + "litellm_teammembership", + "litellm_verificationtoken", + "litellm_organizationtable", + "litellm_tagtable", + "litellm_endusertable", + ): + getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[]) + + @pytest.mark.asyncio async def test_reset_budget_keys_partial_failure(): """ @@ -250,41 +298,18 @@ async def test_reset_budget_users_partial_failure(): @pytest.mark.asyncio -async def test_reset_budget_endusers_partial_failure(): +async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing(): """ - Test that if one enduser fails to reset, the reset loop still processes the other endusers. - We simulate six endsers where the first fails and the others are updated. + A failure anywhere in the budget-tier cascade must persist nothing, so the + tier stays due and the next scheduler tick retries it. Before the fix the + job committed the new budget_reset_at first and zeroed the dependent spend + afterwards, so a failure here left the tier stamped for the next window + while every end user stayed at the cap. """ - user1 = { - "user_id": "user1", - "spend": 20.0, - "budget_id": "budget1", - } # Will trigger simulated failure - user2 = { - "user_id": "user2", - "spend": 25.0, - "budget_id": "budget1", - } # Should be updated - user3 = { - "user_id": "user3", - "spend": 30.0, - "budget_id": "budget1", - } # Should be updated - user4 = { - "user_id": "user4", - "spend": 35.0, - "budget_id": "budget1", - } # Should be updated - user5 = { - "user_id": "user5", - "spend": 40.0, - "budget_id": "budget1", - } # Should be updated - user6 = { - "user_id": "user6", - "spend": 45.0, - "budget_id": "budget1", - } # Should be updated + endusers = [ + _attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"}) + for i in range(1, 7) + ] budget1 = LiteLLM_BudgetTableFull( **{ @@ -301,23 +326,13 @@ async def test_reset_budget_endusers_partial_failure(): if table_name == "budget": return [budget1] elif table_name == "enduser": - return [user1, user2, user3, user4, user5, user6] + return endusers return [] prisma_client.get_data = AsyncMock() prisma_client.get_data.side_effect = get_data_mock - prisma_client.update_data = AsyncMock() - # Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) + batch_calls = _wire_batcher_for_test(prisma_client, fail_commit=True) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -326,41 +341,13 @@ async def test_reset_budget_endusers_partial_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_enduser(enduser): - if enduser["user_id"] == "user1": - raise Exception("Simulated failure for user1") - enduser["spend"] = 0.0 - return enduser + await job.reset_budget_for_litellm_budget_table() + await asyncio.sleep(0.1) - async def fake_reset_team_members(budgets_to_reset): - return 1 - - with ( - patch.object( - ResetBudgetJob, - "_reset_budget_for_enduser", - side_effect=fake_reset_enduser, - ) as mock_reset_enduser, - patch.object( - ResetBudgetJob, - "reset_budget_for_litellm_team_members", - side_effect=fake_reset_team_members, - ) as mock_reset_team_members, - ): - await job.reset_budget_for_litellm_budget_table() - await asyncio.sleep(0.1) - - assert mock_reset_enduser.call_count == 6 - assert prisma_client.update_data.await_count == 2 - update_call = prisma_client.update_data.call_args - assert update_call.kwargs.get("table_name") == "enduser" - updated_users = update_call.kwargs.get("data_list", []) - assert len(updated_users) == 5 - assert updated_users[0]["user_id"] == "user2" - assert updated_users[1]["user_id"] == "user3" - assert updated_users[2]["user_id"] == "user4" - assert updated_users[3]["user_id"] == "user5" - assert updated_users[4]["user_id"] == "user6" + assert batch_calls == [], "a failed cascade must not persist any write" + assert ( + prisma_client.update_data.await_count == 0 + ), "budget_reset_at must not be advanced outside the cascade transaction" failure_hook_calls = ( proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list @@ -369,6 +356,66 @@ async def test_reset_budget_endusers_partial_failure(): call.kwargs.get("call_type") == "reset_budget_endusers" for call in failure_hook_calls ) + proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called() + + +@pytest.mark.asyncio +async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance(): + """ + The happy path: every end user the tier gates is zeroed and the tier's + budget_reset_at advances, all inside one transaction. + """ + endusers = [ + _attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"}) + for i in range(1, 7) + ] + + budget1 = LiteLLM_BudgetTableFull( + **{ + "budget_id": "budget1", + "max_budget": 65.0, + "budget_duration": "2d", + "created_at": datetime.now(timezone.utc) - timedelta(days=3), + } + ) + + prisma_client = MagicMock() + + async def get_data_mock(table_name, *args, **kwargs): + if table_name == "budget": + return [budget1] + elif table_name == "enduser": + return endusers + return [] + + prisma_client.get_data = AsyncMock() + prisma_client.get_data.side_effect = get_data_mock + prisma_client.update_data = AsyncMock() + batch_calls = _wire_batcher_for_test(prisma_client) + + proxy_logging_obj = MagicMock() + proxy_logging_obj.service_logging_obj = MagicMock() + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + + job = ResetBudgetJob(proxy_logging_obj, prisma_client) + + await job.reset_budget_for_litellm_budget_table() + await asyncio.sleep(0.1) + + assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction" + + enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] + assert len(enduser_writes) == 1 + assert enduser_writes[0]["where"]["user_id"]["in"] == [f"user{i}" for i in range(1, 7)] + assert enduser_writes[0]["data"] == {"spend": 0} + + budget_writes = [c for c in batch_calls if c["table"] == "budget"] + assert len(budget_writes) == 1 + assert budget_writes[0]["where"] == {"budget_id": "budget1"} + assert budget_writes[0]["data"]["budget_reset_at"] > datetime.now(timezone.utc) + + proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called() @pytest.mark.asyncio @@ -500,16 +547,8 @@ async def test_reset_budget_continues_other_categories_on_failure(): key1, key2 = _attrify(key1), _attrify(key2) user1, user2 = _attrify(user1), _attrify(user2) team1, team2 = _attrify(team1), _attrify(team2) - # Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) + enduser1 = _attrify(enduser1) + _wire_cascade_reads_for_test(prisma_client) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -541,13 +580,6 @@ async def test_reset_budget_continues_other_categories_on_failure(): ).isoformat() return team - async def fake_reset_enduser(enduser): - enduser["spend"] = 0.0 - return enduser - - async def fake_reset_team_members(budgets_to_reset): - return 1 - with ( patch.object( ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key @@ -558,14 +590,6 @@ async def test_reset_budget_continues_other_categories_on_failure(): patch.object( ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team ) as mock_reset_team, - patch.object( - ResetBudgetJob, "_reset_budget_for_enduser", side_effect=fake_reset_enduser - ) as mock_reset_enduser, - patch.object( - ResetBudgetJob, - "reset_budget_for_litellm_team_members", - side_effect=fake_reset_team_members, - ) as mock_reset_team_members, ): # Call the overall reset_budget method. await job.reset_budget() @@ -575,29 +599,22 @@ async def test_reset_budget_continues_other_categories_on_failure(): called_tables = { call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list } - if mock_reset_team_members.call_count > 0: - called_tables.add("team_membership") - assert called_tables == { - "key", - "user", - "team", - "budget", - "enduser", - "team_membership", - } + assert called_tables == {"key", "user", "team", "budget", "enduser"} - # After the fix, keys/users/teams write via prisma.db.batch_().
.update, - # so only budget + enduser still go through update_data. - calls = prisma_client.update_data.await_args_list - update_data_tables = [c.kwargs.get("table_name") for c in calls] - assert sorted(update_data_tables) == ["budget", "enduser"] + # Every category writes through the batch path now, so update_data is unused. + prisma_client.update_data.assert_not_awaited() - # Check enduser update: enduser succeed. - enduser_call = next(c for c in calls if c.kwargs.get("table_name") == "enduser") - assert len(enduser_call.kwargs.get("data_list", [])) == 1 + # The budget tier's cascade still ran despite the failing user category. + assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1 + enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] + assert len(enduser_writes) == 1 + assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1"]}} + assert enduser_writes[0]["data"] == {"spend": 0} # Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams. - key_writes = [c for c in batch_calls if c["table"] == "key"] + # `op` separates the per-row resets from the cascade sweep, which also + # targets the key table. + key_writes = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"] user_writes = [c for c in batch_calls if c["table"] == "user"] team_writes = [c for c in batch_calls if c["table"] == "team"] assert len(key_writes) == 2 @@ -974,12 +991,12 @@ async def test_service_logger_teams_failure(): @pytest.mark.asyncio async def test_service_logger_endusers_success(): """ - Test that when resetting endusers succeeds the service logger success hook is called with - the correct metadata and no exception is logged. + Test that when the budget-tier cascade commits, the service logger success + hook is called with the correct metadata and no exception is logged. """ endusers = [ - {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}, - {"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}, + _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}), + _attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}), ] budgets = [ LiteLLM_BudgetTableFull( @@ -1002,16 +1019,8 @@ async def test_service_logger_endusers_success(): prisma_client = MagicMock() prisma_client.get_data = AsyncMock(side_effect=fake_get_data) prisma_client.update_data = AsyncMock() - # Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) + batch_calls = _wire_batcher_for_test(prisma_client) + _wire_cascade_reads_for_test(prisma_client) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -1020,31 +1029,16 @@ async def test_service_logger_endusers_success(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_enduser(enduser): - enduser["spend"] = 0.0 - return enduser + with patch( + "litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception" + ) as mock_verbose_exc: + await job.reset_budget_for_litellm_budget_table() + await asyncio.sleep(0.1) + mock_verbose_exc.assert_not_called() - async def fake_reset_team_members(budgets_to_reset): - return 1 - - with ( - patch.object( - ResetBudgetJob, - "_reset_budget_for_enduser", - side_effect=fake_reset_enduser, - ) as mock_reset_enduser, - patch.object( - ResetBudgetJob, - "reset_budget_for_litellm_team_members", - side_effect=fake_reset_team_members, - ) as mock_reset_team_members, - ): - with patch( - "litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception" - ) as mock_verbose_exc: - await job.reset_budget_for_litellm_budget_table() - await asyncio.sleep(0.1) - mock_verbose_exc.assert_not_called() + enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] + assert len(enduser_writes) == 1 + assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1", "user2"]}} proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once() ( @@ -1062,12 +1056,12 @@ async def test_service_logger_endusers_success(): @pytest.mark.asyncio async def test_service_logger_endusers_failure(): """ - Test that a failure during enduser reset calls the failure hook with appropriate metadata, - logs the exception, and does not call the success hook. + Test that a failed cascade calls the failure hook with the rows it had + found, logs the exception, and does not call the success hook. """ endusers = [ - {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}, - {"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}, + _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}), + _attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}), ] budgets = [ LiteLLM_BudgetTableFull( @@ -1090,16 +1084,8 @@ async def test_service_logger_endusers_failure(): prisma_client = MagicMock() prisma_client.get_data = AsyncMock(side_effect=fake_get_data) prisma_client.update_data = AsyncMock() - # Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 0} - ) - # Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) + _wire_batcher_for_test(prisma_client, fail_commit=True) + _wire_cascade_reads_for_test(prisma_client) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -1108,39 +1094,16 @@ async def test_service_logger_endusers_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_enduser(enduser): - if enduser["user_id"] == "user1": - raise Exception("Simulated failure for user1") - enduser["spend"] = 0.0 - return enduser - - async def fake_reset_team_members(budgets_to_reset): - return 1 - - with ( - patch.object( - ResetBudgetJob, - "_reset_budget_for_enduser", - side_effect=fake_reset_enduser, - ) as mock_reset_enduser, - patch.object( - ResetBudgetJob, - "reset_budget_for_litellm_team_members", - side_effect=fake_reset_team_members, - ) as mock_reset_team_members, - ): - with patch( - "litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception" - ) as mock_verbose_exc: - await job.reset_budget_for_litellm_budget_table() - await asyncio.sleep(0.1) - # Verify exception logging - assert mock_verbose_exc.call_count >= 1 - # Verify exception was logged with correct message - assert any( - "Failed to reset budget for enduser" in str(call.args) - for call in mock_verbose_exc.call_args_list - ) + with patch( + "litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception" + ) as mock_verbose_exc: + await job.reset_budget_for_litellm_budget_table() + await asyncio.sleep(0.1) + # The log must name the whole cascade, not just end users: the write + # that failed could have been any of team member / enduser / org / tag + # spend or the budget_reset_at advance. + assert mock_verbose_exc.call_count == 1 + assert "budget table cascade" in str(mock_verbose_exc.call_args.args[0]) proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once() ( @@ -1158,8 +1121,8 @@ async def test_service_logger_endusers_failure(): @pytest.mark.asyncio async def test_reset_budget_for_litellm_team_members_called(): """ - Test that when reset_budget_for_litellm_budget_table is called, - team members' budgets are also reset via reset_budget_for_litellm_team_members + Test that when reset_budget_for_litellm_budget_table is called, team + members' spend is zeroed as part of the cascade transaction. """ # Arrange budget1 = LiteLLM_BudgetTableFull( @@ -1171,7 +1134,7 @@ async def test_reset_budget_for_litellm_team_members_called(): } ) - enduser1 = {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"} + enduser1 = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}) prisma_client = MagicMock() @@ -1184,20 +1147,9 @@ async def test_reset_budget_for_litellm_team_members_called(): prisma_client.get_data = AsyncMock(side_effect=fake_get_data) prisma_client.update_data = AsyncMock() - - # Mock the db.litellm_teammembership.update_many call prisma_client.db = MagicMock() - prisma_client.db.litellm_teammembership = MagicMock() - prisma_client.db.litellm_teammembership.update_many = AsyncMock( - return_value={"count": 2} - ) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 0} - ) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 0} - ) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0}) + batch_calls = _wire_batcher_for_test(prisma_client) + _wire_cascade_reads_for_test(prisma_client) proxy_logging_obj = MagicMock() proxy_logging_obj.service_logging_obj = MagicMock() @@ -1206,23 +1158,11 @@ async def test_reset_budget_for_litellm_team_members_called(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_enduser(enduser): - enduser["spend"] = 0.0 - return enduser - - with patch.object( - ResetBudgetJob, - "_reset_budget_for_enduser", - side_effect=fake_reset_enduser, - ): - # Act - await job.reset_budget_for_litellm_budget_table() + # Act + await job.reset_budget_for_litellm_budget_table() # Assert - # Verify that the team membership update was called - prisma_client.db.litellm_teammembership.update_many.assert_called_once() - - # Verify the call was made with correct parameters - call_args = prisma_client.db.litellm_teammembership.update_many.call_args - assert call_args.kwargs["where"]["budget_id"]["in"] == ["budget1"] - assert call_args.kwargs["data"]["spend"] == 0 + team_member_writes = [c for c in batch_calls if c["table"] == "team_membership"] + assert len(team_member_writes) == 1 + assert team_member_writes[0]["where"]["budget_id"]["in"] == ["budget1"] + assert team_member_writes[0]["data"] == {"spend": 0} diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 616ad8a0981..608dc8cb5c8 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -2,7 +2,6 @@ import asyncio import json import os import sys -import time import types from datetime import datetime, timedelta, timezone from datetime import time as dt_time @@ -13,33 +12,19 @@ import pytest sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path -from litellm._logging import verbose_proxy_logger from litellm.proxy._types import LiteLLM_VerificationToken +from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings -from litellm.proxy.utils import ProxyLogging # Mock classes for testing -class MockLiteLLMTeamMembership: - async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: - # Mock the update_many method for litellm_teammembership - return {"count": 1} +class MockTable: + """A single prisma table: records reads/writes and replays canned rows.""" - -class MockLiteLLMVerificationToken: def __init__(self): - self.update_many_calls: List[Dict[str, Any]] = [] - - async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: - self.update_many_calls.append({"where": where, "data": data}) - return {"count": 1} - - -class MockLiteLLMOrganizationTable: - def __init__(self): - self.update_many_calls: List[Dict[str, Any]] = [] self.find_many_calls: List[Dict[str, Any]] = [] + self.update_many_calls: List[Dict[str, Any]] = [] self._find_many_results: List[Any] = [] def set_find_many_results(self, results: List[Any]): @@ -54,43 +39,12 @@ class MockLiteLLMOrganizationTable: return {"count": 1} -class MockLiteLLMTagTable: - def __init__(self): - self.update_many_calls: List[Dict[str, Any]] = [] - self.find_many_calls: List[Dict[str, Any]] = [] - self._find_many_results: List[Any] = [] - - def set_find_many_results(self, results: List[Any]): - self._find_many_results = results - - async def find_many(self, where: Dict[str, Any]) -> List[Any]: - self.find_many_calls.append({"where": where}) - return self._find_many_results - - async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: - self.update_many_calls.append({"where": where, "data": data}) - return {"count": 1} - - -class MockLiteLLMEndUserTable: - def __init__(self): - self.find_many_calls: List[Dict[str, Any]] = [] - self._find_many_results: List[Any] = [] - - def set_find_many_results(self, results: List[Any]): - self._find_many_results = results - - async def find_many(self, where: Dict[str, Any]) -> List[Any]: - self.find_many_calls.append({"where": where}) - return self._find_many_results - - class MockBatcher: - """Captures per-row update calls and exposes them after commit(). + """Captures the writes queued on one `db.batch_()` and whether it committed. - Mirrors prisma's `db.batch_()` ergonomics enough that the reset job's - narrow-write helpers (`_write_key_reset_updates` et al) can run against - the mock and the test can assert on what would have been written. + Mirrors prisma's batch ergonomics enough that the reset job's write helpers + can run against the mock, and keeps `committed` so tests can prove a failed + cascade persisted nothing. """ def __init__(self): @@ -102,12 +56,23 @@ class MockBatcher: _self._table_name = table_name _self._outer = outer + def _record(_self, op, where, data): + _self._outer.calls.append({"table": _self._table_name, "op": op, "where": where, "data": data}) + def update(_self, where, data): - _self._outer.calls.append({"table": _self._table_name, "where": where, "data": data}) + _self._record("update", where, data) + + def update_many(_self, where, data): + _self._record("update_many", where, data) self.litellm_verificationtoken = _Table("key", self) self.litellm_usertable = _Table("user", self) self.litellm_teamtable = _Table("team", self) + self.litellm_budgettable = _Table("budget", self) + self.litellm_teammembership = _Table("team_membership", self) + self.litellm_organizationtable = _Table("org", self) + self.litellm_tagtable = _Table("tag", self) + self.litellm_endusertable = _Table("enduser", self) async def commit(self): self.committed = True @@ -116,16 +81,20 @@ class MockBatcher: class MockDB: def __init__(self): - self.litellm_teammembership = MockLiteLLMTeamMembership() - self.litellm_verificationtoken = MockLiteLLMVerificationToken() - self.litellm_endusertable = MockLiteLLMEndUserTable() - self.litellm_organizationtable = MockLiteLLMOrganizationTable() - self.litellm_tagtable = MockLiteLLMTagTable() + self.litellm_teammembership = MockTable() + self.litellm_verificationtoken = MockTable() + self.litellm_endusertable = MockTable() + self.litellm_organizationtable = MockTable() + self.litellm_tagtable = MockTable() self.batch_calls: List[Dict[str, Any]] = [] + self.batchers: List[MockBatcher] = [] def batch_(self): batcher = MockBatcher() - # Aggregate calls across all batches so tests can assert on cumulative writes. + self.batchers.append(batcher) + # Aggregate calls across all batches so tests can assert on cumulative + # writes. Only committed batches contribute: an abandoned batch writes + # nothing, exactly as prisma behaves. original_commit = batcher.commit async def _record_and_commit(): @@ -152,9 +121,11 @@ class MockPrismaClient: "budget": [], "enduser": [], } + self.get_data_calls: List[Dict[str, Any]] = [] self.db = MockDB() async def get_data(self, table_name, query_type, **kwargs): + self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs}) data = self.data.get(table_name, []) # Handle specific filtering for budget table queries @@ -218,6 +189,39 @@ async def run_async_test(coro): return await coro +_ALREADY_EXPIRED = object() + + +def _budget_row( + budget_id: str = "test-budget-1", + budget_duration: Any = "7d", + budget_reset_at: Any = _ALREADY_EXPIRED, + max_budget: float = 10.0, +): + """An expiring budget tier, shaped like the rows get_data() hands back.""" + now = datetime.now(timezone.utc) + return type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": max_budget, + "budget_duration": budget_duration, + "budget_reset_at": (now - timedelta(hours=1) if budget_reset_at is _ALREADY_EXPIRED else budget_reset_at), + "budget_id": budget_id, + "created_at": now - timedelta(days=30), + }, + ) + + +def _batch_writes(mock_prisma_client, table: str, op: str | None = None) -> List[Dict[str, Any]]: + """Writes that were committed to the DB, optionally narrowed to one op.""" + return [ + call + for call in mock_prisma_client.db.batch_calls + if call["table"] == table and (op is None or call["op"] == op) + ] + + # Tests def test_write_key_reset_updates_skips_none_token_and_still_writes_the_rest(reset_budget_job, mock_prisma_client): """A key with token=None must be skipped, not queued as where={"token": None}. @@ -234,10 +238,10 @@ def test_write_key_reset_updates_skips_none_token_and_still_writes_the_rest(rese asyncio.run(reset_budget_job._write_key_reset_updates(updated_keys=keys)) - key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"] - assert key_writes == [ + assert _batch_writes(mock_prisma_client, "key") == [ { "table": "key", + "op": "update", "where": {"token": "tok-ok"}, "data": {"spend": 0, "budget_reset_at": reset_at}, } @@ -369,18 +373,9 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client): def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client): - # Setup test data + """End-user spend is zeroed and the tier's window advances, in one batch.""" now = datetime.now(timezone.utc) - test_budget = type( - "LiteLLM_BudgetTable", - (), - { - "max_budget": 500.0, - "budget_duration": "1d", - "budget_reset_at": now, - "budget_id": "test-budget-1", - }, - ) + test_budget = _budget_row(budget_id="test-budget-1", budget_duration="1d", budget_reset_at=now) test_enduser = type( "LiteLLM_EndUserTable", @@ -395,16 +390,22 @@ def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client): mock_prisma_client.data["budget"] = [test_budget] mock_prisma_client.data["enduser"] = [test_enduser] - # Run the test asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - # Verify results - assert len(mock_prisma_client.updated_data["enduser"]) == 1 - assert len(mock_prisma_client.updated_data["budget"]) == 1 - updated_enduser = mock_prisma_client.updated_data["enduser"][0] - updated_budget = mock_prisma_client.updated_data["budget"][0] - assert updated_enduser.spend == 0.0 - assert updated_budget.budget_reset_at > now + assert _batch_writes(mock_prisma_client, "enduser") == [ + { + "table": "enduser", + "op": "update_many", + "where": {"user_id": {"in": ["test-enduser-1"]}}, + "data": {"spend": 0}, + } + ] + + budget_writes = _batch_writes(mock_prisma_client, "budget") + assert len(budget_writes) == 1 + assert budget_writes[0]["where"] == {"budget_id": "test-budget-1"} + assert budget_writes[0]["data"]["budget_reset_at"] > now + assert set(budget_writes[0]["data"].keys()) == {"budget_reset_at"} def test_reset_budget_all(reset_budget_job, mock_prisma_client): @@ -485,190 +486,81 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client): ("user", {"user_id": "uid-all-1"}), ("team", {"team_id": "tid-all-1"}), ]: - writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == table_name] + writes = _batch_writes(mock_prisma_client, table_name, op="update") assert len(writes) == 1, f"expected 1 {table_name} write, got {len(writes)}" assert writes[0]["where"] == where assert writes[0]["data"]["spend"] == 0 assert set(writes[0]["data"].keys()) == {"spend", "budget_reset_at"} - # Enduser + budget rows still go through update_data (not narrowed; different path). - assert len(mock_prisma_client.updated_data["enduser"]) == 1 - assert len(mock_prisma_client.updated_data["budget"]) == 1 - assert mock_prisma_client.updated_data["enduser"][0].spend == 0.0 - - -def test_reset_budget_for_keys_linked_to_budgets(reset_budget_job, mock_prisma_client): - """ - Test that when a budget tier is reset, keys linked to that budget - (via budget_id) that don't have their own budget_duration also get - their spend reset. - - This covers the case where keys were created with budget_id but - budget_duration was not inherited to the key (pre-fix keys). - """ - from litellm.proxy._types import LiteLLM_BudgetTableFull - - now = datetime.now(timezone.utc) - - # Create a budget tier that is due for reset - test_budget = type( - "LiteLLM_BudgetTableFull", - (), + # The budget tier's cascade rides the same batch machinery. + assert _batch_writes(mock_prisma_client, "enduser") == [ { - "max_budget": 10.0, - "budget_duration": "7d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "7d-budget-tier", - "created_at": now - timedelta(days=7), + "table": "enduser", + "op": "update_many", + "where": {"user_id": {"in": ["test-enduser-1"]}}, + "data": {"spend": 0}, + } + ] + assert len(_batch_writes(mock_prisma_client, "budget")) == 1 + + +_LINKED_TABLE_CASES = [ + ("team_membership", {"budget_id": {"in": ["7d-budget-tier"]}}), + ( + "key", + { + "budget_id": {"in": ["7d-budget-tier"]}, + "budget_duration": None, + "spend": {"gt": 0}, }, - ) - - budgets_to_reset = [test_budget] - - # Run the method - asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)) - - # Verify that update_many was called on litellm_verificationtoken - calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls - assert len(calls) == 1, f"Expected 1 update_many call, got {len(calls)}" - - # Verify the where clause filters by budget_id and null budget_duration - call = calls[0] - assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]} - assert call["where"]["budget_duration"] is None - - # Verify spend is reset to 0 - assert call["data"]["spend"] == 0 + ), + ("org", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), + ("tag", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), +] -def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_duration( - reset_budget_job, mock_prisma_client +@pytest.mark.parametrize( + "table, expected_where", + _LINKED_TABLE_CASES, + ids=[case[0] for case in _LINKED_TABLE_CASES], +) +def test_budget_table_reset_zeroes_spend_on_every_linked_table( + reset_budget_job, mock_prisma_client, table, expected_where ): + """One expiring tier zeroes spend on every row it gates. + + The filters carry real behavior: keys must be narrowed to + `budget_duration: None` so keys with their own reset schedule aren't + double-reset by reset_budget_for_litellm_keys(), and the payload must stay + exactly {"spend": 0} because `total_spend` is a lifetime counter a reset + may never touch. """ - Test that keys with BOTH budget_id AND budget_duration are excluded from - reset_budget_for_keys_linked_to_budgets. Such keys have their own reset - schedule and are handled only by reset_budget_for_litellm_keys(). The - budget_duration=None filter ensures they are NOT double-reset when the - linked budget tier expires. - """ - now = datetime.now(timezone.utc) + mock_prisma_client.data["budget"] = [_budget_row(budget_id="7d-budget-tier", budget_duration="7d")] - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 10.0, - "budget_duration": "7d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "7d-budget-tier", - "created_at": now - timedelta(days=7), - }, - ) + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - budgets_to_reset = [test_budget] - - asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)) - - calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls - assert len(calls) == 1 - call = calls[0] - - # Critical: budget_duration must be None so keys with their own budget_duration - # (e.g. key has budget_id="X" AND budget_duration=60) are excluded. - # Those keys are reset only by reset_budget_for_litellm_keys() - no double-reset. - assert call["where"]["budget_duration"] is None - assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]} + writes = _batch_writes(mock_prisma_client, table, op="update_many") + assert len(writes) == 1, f"expected exactly 1 {table} write, got {writes}" + assert writes[0]["where"] == expected_where + assert writes[0]["data"] == {"spend": 0} -def test_reset_budget_for_keys_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): - """ - Test that when there are no budgets to reset, no update is performed - on the verification token table. - """ - # Run with empty list - asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=[])) +def test_budget_table_reset_writes_nothing_when_no_budget_is_due(reset_budget_job, mock_prisma_client): + """Nothing due means no transaction is opened at all.""" + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - # Verify no update_many calls were made - calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls - assert len(calls) == 0 + assert mock_prisma_client.db.batchers == [] + assert mock_prisma_client.db.batch_calls == [] -def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_client): - """ - Test that when a budget tier is reset, orgs linked to that budget - (via budget_id) also get their spend reset. - """ - now = datetime.now(timezone.utc) +def _run_reset_at_fixed_now(job, fixed_now): + """Run the budget-table reset with `now` pinned for reset-time math.""" + from unittest.mock import patch - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 100.0, - "budget_duration": "30d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "30d-org-budget", - "created_at": now - timedelta(days=30), - }, - ) - - asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[test_budget])) - - calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls - assert len(calls) == 1 - call = calls[0] - assert call["where"]["budget_id"] == {"in": ["30d-org-budget"]} - assert call["where"]["spend"] == {"gt": 0} - assert call["data"]["spend"] == 0 - - -def test_reset_budget_for_orgs_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): - """ - Test that when there are no budgets to reset, no update is performed - on the organization table. - """ - asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[])) - calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls - assert len(calls) == 0 - - -def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_client): - """ - Test that when a budget tier is reset, tags linked to that budget - (via budget_id) also get their spend reset. - """ - now = datetime.now(timezone.utc) - - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 50.0, - "budget_duration": "30d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "30d-tag-budget", - "created_at": now - timedelta(days=30), - }, - ) - - asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[test_budget])) - - calls = mock_prisma_client.db.litellm_tagtable.update_many_calls - assert len(calls) == 1 - call = calls[0] - assert call["where"]["budget_id"] == {"in": ["30d-tag-budget"]} - assert call["where"]["spend"] == {"gt": 0} - assert call["data"]["spend"] == 0 - - -def test_reset_budget_for_tags_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): - """ - Test that when there are no budgets to reset, no update is performed - on the tag table. - """ - asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[])) - calls = mock_prisma_client.db.litellm_tagtable.update_many_calls - assert len(calls) == 0 + with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: + mock_dt.now.return_value = fixed_now + mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) + asyncio.run(job.reset_budget_for_litellm_budget_table()) @pytest.mark.parametrize( @@ -680,215 +572,70 @@ def test_reset_budget_for_tags_linked_to_budgets_empty(reset_budget_job, mock_pr ], ids=["30d-calendar-month", "1mo-calendar-month", "1d-next-midnight"], ) -def test_reset_budget_reset_at_date_calendar_aligned(budget_duration, expected_day, expected_month): - """ - Verify that _reset_budget_reset_at_date produces calendar-aligned reset - times (matching get_budget_reset_time), not sliding-window offsets. - """ - from unittest.mock import patch - - # Fix "now" to 2023-06-15 10:30:00 UTC for deterministic results +def test_budget_reset_at_written_is_calendar_aligned( + reset_budget_job, mock_prisma_client, budget_duration, expected_day, expected_month +): + """The advanced budget_reset_at is calendar-aligned, not a sliding + now + duration offset.""" fixed_now = datetime(2023, 6, 15, 10, 30, 0, tzinfo=timezone.utc) + mock_prisma_client.data["budget"] = [ + _budget_row( + budget_id="test-budget", + budget_duration=budget_duration, + budget_reset_at=fixed_now - timedelta(hours=1), + ) + ] - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "budget_duration": budget_duration, - "budget_reset_at": fixed_now - timedelta(hours=1), - "budget_id": "test-budget", - "created_at": fixed_now - timedelta(days=30), - }, - ) + _run_reset_at_fixed_now(reset_budget_job, fixed_now) - with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: - mock_dt.now.return_value = fixed_now - mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) - - assert test_budget.budget_reset_at.day == expected_day - assert test_budget.budget_reset_at.month == expected_month - assert test_budget.budget_reset_at.hour == 0 - assert test_budget.budget_reset_at.minute == 0 - assert test_budget.budget_reset_at.second == 0 + writes = _batch_writes(mock_prisma_client, "budget") + assert len(writes) == 1 + written = writes[0]["data"]["budget_reset_at"] + assert (written.day, written.month) == (expected_day, expected_month) + assert (written.hour, written.minute, written.second) == (0, 0, 0) -def test_reset_budget_reset_at_date_7d_next_monday(): - """Verify 7d budget duration resets to next Monday at midnight.""" - from unittest.mock import patch - +def test_budget_reset_at_written_for_7d_is_next_monday(reset_budget_job, mock_prisma_client): + """7d budgets advance to next Monday at midnight.""" # 2023-06-14 is a Wednesday fixed_now = datetime(2023, 6, 14, 10, 30, 0, tzinfo=timezone.utc) + mock_prisma_client.data["budget"] = [ + _budget_row(budget_id="test-budget", budget_duration="7d", budget_reset_at=fixed_now - timedelta(hours=1)) + ] - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "budget_duration": "7d", - "budget_reset_at": fixed_now - timedelta(hours=1), - "budget_id": "test-budget", - "created_at": fixed_now - timedelta(days=7), - }, - ) + _run_reset_at_fixed_now(reset_budget_job, fixed_now) - with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: - mock_dt.now.return_value = fixed_now - mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) - - # Next Monday after Wednesday June 14 is June 19 - assert test_budget.budget_reset_at.day == 19 - assert test_budget.budget_reset_at.month == 6 - assert test_budget.budget_reset_at.weekday() == 0 # Monday - assert test_budget.budget_reset_at.hour == 0 + written = _batch_writes(mock_prisma_client, "budget")[0]["data"]["budget_reset_at"] + assert (written.day, written.month) == (19, 6) + assert written.weekday() == 0 + assert written.hour == 0 -def test_reset_budget_reset_at_date_none_duration(): - """Verify that budget_reset_at is unchanged when budget_duration is None.""" - original_reset_at = datetime(2023, 6, 20, 0, 0, 0, tzinfo=timezone.utc) - now = datetime(2023, 6, 15, 10, 0, 0, tzinfo=timezone.utc) +def test_budget_with_no_duration_gets_no_reset_at_write(reset_budget_job, mock_prisma_client): + """A tier without a duration has no next window, so its row is left alone + rather than rewritten with an unchanged value.""" + mock_prisma_client.data["budget"] = [ + _budget_row( + budget_id="no-duration", budget_duration=None, budget_reset_at=datetime(2023, 6, 20, tzinfo=timezone.utc) + ) + ] - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "budget_duration": None, - "budget_reset_at": original_reset_at, - "budget_id": "test-budget", - "created_at": now - timedelta(days=30), - }, - ) + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, now, BudgetResetSettings())) - assert test_budget.budget_reset_at == original_reset_at + assert _batch_writes(mock_prisma_client, "budget") == [] -def test_reset_budget_reset_at_date_none_reset_at(): - """Verify that budget_reset_at is set correctly even when previously None.""" - from unittest.mock import patch - +def test_budget_reset_at_written_when_previously_null(reset_budget_job, mock_prisma_client): + """A tier whose budget_reset_at was never initialized still gets one.""" fixed_now = datetime(2023, 6, 15, 10, 30, 0, tzinfo=timezone.utc) + mock_prisma_client.data["budget"] = [ + _budget_row(budget_id="test-budget", budget_duration="30d", budget_reset_at=None) + ] - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "budget_duration": "30d", - "budget_reset_at": None, - "budget_id": "test-budget", - "created_at": fixed_now - timedelta(days=5), - }, - ) + _run_reset_at_fixed_now(reset_budget_job, fixed_now) - with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: - mock_dt.now.return_value = fixed_now - mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) - - # Should be set to 1st of next month (July 1) - assert test_budget.budget_reset_at is not None - assert test_budget.budget_reset_at.day == 1 - assert test_budget.budget_reset_at.month == 7 - - -def test_budget_table_reset_also_resets_linked_keys(reset_budget_job, mock_prisma_client): - """ - Integration-style test: when reset_budget_for_litellm_budget_table runs, - it should also reset spend for keys linked to the expiring budget tiers - (in addition to end-users and team members). - """ - now = datetime.now(timezone.utc) - - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 10.0, - "budget_duration": "7d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "7d-budget-tier", - "created_at": now - timedelta(days=7), - }, - ) - - mock_prisma_client.data["budget"] = [test_budget] - - # Run the full budget table reset - asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - - # Verify that keys linked to the budget were also reset - calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls - assert len(calls) == 1, ( - "Expected reset_budget_for_litellm_budget_table to also reset keys " - f"linked to expiring budgets, but got {len(calls)} update_many calls" - ) - assert calls[0]["where"]["budget_id"] == {"in": ["7d-budget-tier"]} - assert calls[0]["data"]["spend"] == 0 - - -def test_budget_table_reset_also_resets_linked_orgs(reset_budget_job, mock_prisma_client): - """ - Integration-style test: when reset_budget_for_litellm_budget_table runs, - it should also reset spend for orgs linked to the expiring budget tiers - (in addition to end-users, team members, and keys). - """ - now = datetime.now(timezone.utc) - - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 100.0, - "budget_duration": "30d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "30d-org-budget", - "created_at": now - timedelta(days=30), - }, - ) - - mock_prisma_client.data["budget"] = [test_budget] - - asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - - calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls - assert len(calls) == 1, ( - "Expected reset_budget_for_litellm_budget_table to also reset orgs " - f"linked to expiring budgets, but got {len(calls)} update_many calls" - ) - assert calls[0]["where"]["budget_id"] == {"in": ["30d-org-budget"]} - assert calls[0]["data"]["spend"] == 0 - - -def test_budget_table_reset_also_resets_linked_tags(reset_budget_job, mock_prisma_client): - """ - Integration-style test: when reset_budget_for_litellm_budget_table runs, - it should also reset spend for tags linked to the expiring budget tiers. - """ - now = datetime.now(timezone.utc) - - test_budget = type( - "LiteLLM_BudgetTableFull", - (), - { - "max_budget": 50.0, - "budget_duration": "30d", - "budget_reset_at": now - timedelta(hours=1), - "budget_id": "30d-tag-budget", - "created_at": now - timedelta(days=30), - }, - ) - - mock_prisma_client.data["budget"] = [test_budget] - - asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - - calls = mock_prisma_client.db.litellm_tagtable.update_many_calls - assert len(calls) == 1, ( - "Expected reset_budget_for_litellm_budget_table to also reset tags " - f"linked to expiring budgets, but got {len(calls)} update_many calls" - ) - assert calls[0]["where"]["budget_id"] == {"in": ["30d-tag-budget"]} - assert calls[0]["data"]["spend"] == 0 + written = _batch_writes(mock_prisma_client, "budget")[0]["data"]["budget_reset_at"] + assert (written.day, written.month) == (1, 7) def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock_prisma_client): @@ -965,16 +712,14 @@ def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - # Both end users should have been reset - updated = mock_prisma_client.updated_data["enduser"] - assert len(updated) == 2, f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}" - - user_ids = {u.user_id for u in updated} - assert "enduser-explicit" in user_ids - assert "enduser-implicit" in user_ids - - for u in updated: - assert u.spend == 0.0, f"Expected spend=0 for {u.user_id}, got {u.spend}" + # Both end users are zeroed by the same committed statement. + enduser_writes = _batch_writes(mock_prisma_client, "enduser") + assert len(enduser_writes) == 1, f"Expected a single enduser write, got {enduser_writes}" + assert set(enduser_writes[0]["where"]["user_id"]["in"]) == { + "enduser-explicit", + "enduser-implicit", + } + assert enduser_writes[0]["data"] == {"spend": 0} # Verify find_many was called to fetch NULL-budget-id end users find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls @@ -1054,34 +799,6 @@ def test_reset_budget_skips_null_budget_id_endusers_when_default_not_in_reset_li litellm.max_end_user_budget_id = None -def test_reset_budget_for_team_members_preserves_total_spend(): - """Regression guard: reset_budget_for_litellm_team_members must zero `spend` - but leave `total_spend` untouched. - - The reset writes `data={"spend": 0}` explicitly. If a future refactor adds - `"total_spend": 0` to that dict, this test fails immediately. - """ - expired_budget = type( - "LiteLLM_BudgetTableFull", - (), - {"budget_id": "budget-1"}, - ) - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[]) - mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client) - - asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) - - mock_prisma_client.db.litellm_teammembership.update_many.assert_called_once() - call_kwargs = mock_prisma_client.db.litellm_teammembership.update_many.call_args.kwargs - assert call_kwargs["where"]["budget_id"]["in"] == ["budget-1"] - assert call_kwargs["data"] == {"spend": 0} - assert "total_spend" not in call_kwargs["data"] - - # --------------------------------------------------------------------------- # reset_budget_windows (per-key / per-team concurrent window resets) # --------------------------------------------------------------------------- @@ -1323,28 +1040,6 @@ def _make_counter_invalidation_job(monkeypatch): return spend_counter_cache -def test_reset_budget_for_team_members_invalidates_redis_counter(monkeypatch): - """Team-member budget reset clears the Redis spend counter.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - membership = type( - "Membership", - (), - {"user_id": "alice", "team_id": "team-x", "budget_id": "budget-1"}, - ) - - prisma_client = MagicMock() - prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership]) - prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) - - counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team_member:alice:team-x", value=0.0, ttl=60) - counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:team_member:alice:team-x", value=0.0, ttl=60) - - def test_reset_budget_for_keys_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch): """Key budget reset must clear the Redis spend counter.""" counter_cache = _make_counter_invalidation_job(monkeypatch) @@ -1574,207 +1269,240 @@ def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, ) -def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monkeypatch): - """Resetting keys via budget tier must clear each linked key's counter.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_key = type("Key", (), {"token": "sk-linked"}) - - prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key]) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget])) - - counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-linked", value=0.0, ttl=60) +_INVALIDATION_CASES = [ + ( + "litellm_teammembership", + type("Membership", (), {"user_id": "alice", "team_id": "team-x", "budget_id": "budget-1"}), + "spend:team_member:alice:team-x", + {"team-x_alice"}, + ), + ( + "litellm_verificationtoken", + type("Key", (), {"token": "sk-linked"}), + "spend:key:sk-linked", + {"sk-linked"}, + ), + ( + "litellm_organizationtable", + type("Org", (), {"organization_id": "org-acme"}), + "spend:org:org-acme", + {"org_id:org-acme", "org_id:org-acme:with_budget"}, + ), + ( + "litellm_tagtable", + type("Tag", (), {"tag_name": "tenant-42"}), + "spend:tag:tenant-42", + {"tag:tenant-42"}, + ), +] -def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monkeypatch): - """Resetting orgs via budget tier must clear each linked org's counter.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_org = type("Org", (), {"organization_id": "org-acme"}) - - prisma_client = MagicMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org]) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget])) - - counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:org:org-acme", value=0.0, ttl=60) - counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:org:org-acme", value=0.0, ttl=60) - - -def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monkeypatch): - """Resetting tags via budget tier must clear each linked tag's counter.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) - - prisma_client = MagicMock() - prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - - counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:tag:tenant-42", value=0.0, ttl=60) - counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:tag:tenant-42", value=0.0, ttl=60) - - -def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache( - monkeypatch, +@pytest.mark.parametrize( + "table_attr, linked_row, counter_key, cache_keys", + _INVALIDATION_CASES, + ids=["team_membership", "key", "org", "tag"], +) +def test_budget_table_reset_invalidates_counters_and_management_cache( + reset_budget_job, mock_prisma_client, monkeypatch, table_attr, linked_row, counter_key, cache_keys ): - """Regression guard for the bug where tag spend stayed frozen across cycles. + """Every row the cascade zeroes gets its spend counter cleared and its + management-cache entry dropped. - ``SpendCounterReseed.from_db`` returns ``None`` for ``spend:tag:*`` keys, - so once the spend counter expires the tag budget check falls back to the - cached ``LiteLLM_TagTable.spend``. If we don't drop the management cache - entry on reset, that cached object lingers (TTL 60s) with the pre-reset - spend, and ``_tag_max_budget_check`` keeps returning HTTP 400 even though - the DB row has been zeroed. + Both matter. ``SpendCounterReseed.from_db`` returns None for tags, so once + the counter expires the budget check falls back to the cached row's + ``.spend``; and for keys, orgs and team memberships another pod's cached + object can stay pinned above the zeroed DB row until its TTL. Team + membership cache keys follow auth's ``{team_id}_{user_id}`` shape, and orgs + carry both the plain and the ``:with_budget`` entry. """ counter_cache = _make_counter_invalidation_job(monkeypatch) + mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-1")] + getattr(mock_prisma_client.db, table_attr).set_find_many_results([linked_row]) - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - prisma_client = MagicMock() - prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="tag:tenant-42") + counter_cache.in_memory_cache.set_cache.assert_any_call(key=counter_key, value=0.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key=counter_key, value=0.0, ttl=60) + deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} + assert cache_keys <= deleted -def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management_cache( - monkeypatch, -): - """When multiple tags share the expired budget tier, every one of them - has its ``user_api_key_cache`` entry dropped — not just the first.""" +def test_budget_table_reset_invalidates_every_tag_not_just_the_first(reset_budget_job, mock_prisma_client, monkeypatch): + """When several tags share the expiring tier, all of them are evicted.""" counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_tags = [ - type("Tag", (), {"tag_name": "tenant-a"}), - type("Tag", (), {"tag_name": "tenant-b"}), - type("Tag", (), {"tag_name": "tenant-c"}), - ] - - prisma_client = MagicMock() - prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=linked_tags) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 3}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - - deleted_keys = { - call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list - } - assert deleted_keys == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"} - - -def test_reset_budget_for_keys_linked_to_budgets_invalidates_management_cache( - monkeypatch, -): - """Budget-tier key resets must drop the cached key object (hashed token key). - - Historically this test used ``assert_not_awaited()`` on - ``user_api_key_cache.async_delete_cache``, reflecting the assumption that - ``SpendCounterReseed.from_db`` alone kept spend consistent for keys and - that invalidating the management cache was unnecessary. That was flipped to - ``assert_any_await(...)`` because the old invariant fails across pods: a - budget reset on one instance can leave another pod's cached key object - (including embedded ``.spend``) stale until TTL expiry. Eviction now matches - tags/orgs/teams. Do not treat the ``cache_key_fn`` / invalidation wiring as - redundant without revisiting that cross-pod consistency story. - """ - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_key = type("Key", (), {"token": "sk-linked"}) - - prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key]) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget])) - - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="sk-linked") - - -def test_reset_budget_for_orgs_linked_to_budgets_invalidates_management_cache( - monkeypatch, -): - """Org rows use both base and budget-table cache keys — evict both on reset.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_org = type("Org", (), {"organization_id": "org-acme"}) - - prisma_client = MagicMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org]) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1}) - - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget])) - - deleted_keys = { - call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list - } - assert deleted_keys == { - "org_id:org-acme", - "org_id:org-acme:with_budget", - } - - -def test_reset_budget_for_team_members_invalidates_management_cache(monkeypatch): - """Team membership cache key matches auth: ``{team_id}_{user_id}``.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - - expired_budget = type("B", (), {"budget_id": "budget-1"}) - membership = type( - "Membership", - (), - {"user_id": "alice", "team_id": "team-x", "budget_id": "budget-1"}, + mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-1")] + mock_prisma_client.db.litellm_tagtable.set_find_many_results( + [type("Tag", (), {"tag_name": name}) for name in ("tenant-a", "tenant-b", "tenant-c")] ) - prisma_client = MagicMock() - prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership]) - prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) - - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="team-x_alice") + deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} + assert deleted == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"} -def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure_still_resets( - monkeypatch, -): - """If ``async_delete_cache`` raises, the DB cascade must still complete.""" +def test_budget_table_reset_commits_even_when_cache_eviction_fails(reset_budget_job, mock_prisma_client, monkeypatch): + """Eviction runs after the commit, so a broken cache cannot undo the write.""" counter_cache = _make_counter_invalidation_job(monkeypatch) counter_cache.user_api_key_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("cache unavailable")) + mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-1")] + mock_prisma_client.db.litellm_tagtable.set_find_many_results([type("Tag", (), {"tag_name": "tenant-42"})]) - expired_budget = type("B", (), {"budget_id": "budget-1"}) - linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - prisma_client = MagicMock() - prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) - prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) + assert len(_batch_writes(mock_prisma_client, "tag", op="update_many")) == 1 + assert mock_prisma_client.db.batchers[0].committed is True - job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) - asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - prisma_client.db.litellm_tagtable.update_many.assert_awaited_once() +# --------------------------------------------------------------------------- +# Atomicity of the budget-table cascade (LIT-5138) +# --------------------------------------------------------------------------- + + +class FailingCommitDB(MockDB): + """Batches that blow up at commit, like a Postgres timeout mid-cascade.""" + + def batch_(self): + batcher = super().batch_() + + async def _fail(): + raise RuntimeError("simulated Postgres timeout mid-cascade") + + batcher.commit = _fail + return batcher + + +class FailingTeamMembershipDB(MockDB): + """Queueing the team-membership reset raises, i.e. the cascade breaks after + earlier writes are already queued.""" + + def batch_(self): + batcher = super().batch_() + + def _fail(where, data): + raise RuntimeError("simulated failure queueing the team-membership reset") + + batcher.litellm_teammembership.update_many = _fail + return batcher + + +class OrderRecordingDB(MockDB): + """Appends a marker to a shared list when a batch commits.""" + + def __init__(self, events): + super().__init__() + self._events = events + + def batch_(self): + batcher = super().batch_() + wrapped = batcher.commit + + async def _record_commit(): + self._events.append("commit") + return await wrapped() + + batcher.commit = _record_commit + return batcher + + +def _job_with_expired_budget(db, proxy_logging=None): + """A job with one due tier and a linked tag, so cache invalidation has + something to invalidate and its absence is a real signal.""" + prisma_client = MockPrismaClient() + prisma_client.db = db + prisma_client.data["budget"] = [_budget_row(budget_id="budget-1", budget_duration="7d")] + db.litellm_tagtable.set_find_many_results([type("Tag", (), {"tag_name": "tenant-42"})]) + job = ResetBudgetJob( + proxy_logging_obj=proxy_logging or MockProxyLogging(), + prisma_client=prisma_client, + ) + return job, prisma_client + + +@pytest.mark.parametrize( + "db_factory", + [FailingCommitDB, FailingTeamMembershipDB], + ids=["commit-fails", "queueing-fails"], +) +def test_budget_reset_at_is_not_advanced_when_the_cascade_fails(db_factory, monkeypatch): + """Regression for LIT-5138. + + The old code committed the new budget_reset_at first and zeroed the + dependent spend afterwards. A failure part-way through left the tier + stamped for the next window, so every later tick skipped it and team + member / enduser / org / tag spend stayed at the cap for the whole window. + One transaction means a failure anywhere persists nothing and the tier is + still due on the next tick. + """ + counter_cache = _make_counter_invalidation_job(monkeypatch) + job, prisma_client = _job_with_expired_budget(db_factory()) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) # swallowed, retried next tick + + assert prisma_client.db.batch_calls == [], "a failed cascade must not persist any write" + assert prisma_client.db.batchers[0].committed is False + assert prisma_client.updated_data["budget"] == [], "budget_reset_at must not be advanced outside the transaction" + counter_cache.in_memory_cache.set_cache.assert_not_called() + counter_cache.user_api_key_cache.async_delete_cache.assert_not_awaited() + + +def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mock_prisma_client, monkeypatch): + """Dependent spend and the budget_reset_at advance ride one batch.""" + _make_counter_invalidation_job(monkeypatch) + now = datetime.now(timezone.utc) + budget = _budget_row(budget_id="budget-1", budget_duration="7d") + mock_prisma_client.data["budget"] = [budget] + mock_prisma_client.data["enduser"] = [ + type("EndUser", (), {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1"}) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + assert len(mock_prisma_client.db.batchers) == 1, "the cascade must not be split across transactions" + batcher = mock_prisma_client.db.batchers[0] + assert batcher.committed is True + assert {(call["table"], call["op"]) for call in batcher.calls} == { + ("team_membership", "update_many"), + ("key", "update_many"), + ("org", "update_many"), + ("tag", "update_many"), + ("enduser", "update_many"), + ("budget", "update_many"), + } + budget_write = next(call for call in batcher.calls if call["table"] == "budget") + assert budget_write["data"]["budget_reset_at"] > now + + +def test_caches_are_invalidated_only_after_the_transaction_commits(monkeypatch): + """A counter zeroed before the write lands would admit requests past the + cap while the DB still holds the over-budget spend.""" + events = [] + counter_cache = _make_counter_invalidation_job(monkeypatch) + counter_cache.in_memory_cache.set_cache.side_effect = lambda **kwargs: events.append("counter") + + job, _ = _job_with_expired_budget(OrderRecordingDB(events)) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert events == ["commit", "counter"] + + +def test_failed_cascade_is_logged_as_a_cascade_failure(monkeypatch): + """The failure log has to name what actually broke. The old catch-all + blamed end users even when the team-membership write was the failure.""" + from unittest.mock import patch + + _make_counter_invalidation_job(monkeypatch) + job, _ = _job_with_expired_budget(FailingTeamMembershipDB()) + + with patch("litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception") as mock_exception: + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert mock_exception.call_count == 1 + message = mock_exception.call_args.args[0] + assert "cascade" in message + for mentioned in ("team member", "enduser", "org", "tag", "budget_reset_at"): + assert mentioned in message, f"failure log should mention {mentioned}: {message}" def _extract_reset_where(find_many_mock): @@ -1799,23 +1527,65 @@ def _asserts_null_reset_is_due(where): branches = where.get("OR") assert isinstance(branches, list), f"expected an OR filter, got {where!r}" - has_null_branch = any( - b.get("AND") - == [ - {"budget_reset_at": None}, - {"NOT": {"budget_duration": None}}, - ] - for b in branches - if isinstance(b, dict) - ) - has_expired_branch = any( - isinstance(b, dict) - and "budget_reset_at" in b - and b["budget_reset_at"] is not None - for b in branches - ) + has_null_branch = {"budget_reset_at": None} in branches + has_expired_branch = any(isinstance(b, dict) and isinstance(b.get("budget_reset_at"), dict) for b in branches) assert has_null_branch, f"missing NULL-reset_at branch in {where!r}" assert has_expired_branch, f"missing expired-reset_at branch in {where!r}" + assert where.get("NOT") == {"budget_duration": None}, f"NULL reset_at is only due with a duration: {where!r}" + + +_RESET_TABLE_ATTRS = { + "user": "litellm_usertable", + "team": "litellm_teamtable", + "budget": "litellm_budgettable", + "key": "litellm_verificationtoken", +} + + +def _run_reset_query(table_name, **extra): + """Run ``get_data`` for one table's budget-reset query against a mocked + prisma handle, and hand back the ``find_many`` mock it drove.""" + from litellm.proxy.utils import PrismaClient + + client = PrismaClient.__new__(PrismaClient) + client.db = MagicMock() + find_many = AsyncMock(return_value=[]) + setattr(getattr(client.db, _RESET_TABLE_ATTRS[table_name]), "find_many", find_many) + + now = datetime.now(timezone.utc) + expires = {"expires": now} if table_name == "key" else {} + asyncio.run(client.get_data(table_name=table_name, query_type="find_all", reset_at=now, **expires, **extra)) + return find_many + + +@pytest.mark.parametrize("table_name", ["user", "team", "budget", "key"]) +def test_get_data_reset_query_applies_the_row_limit(table_name): + """The reset job pages through due rows, so ``limit`` has to reach prisma as + ``take``. Dropped, every worker goes back to pulling the entire expired set + in one unbounded query at the same calendar boundary.""" + find_many = _run_reset_query(table_name, limit=7) + + assert find_many.await_args.kwargs["take"] == 7 + + +@pytest.mark.parametrize("table_name", ["user", "team", "budget", "key"]) +def test_get_data_reset_query_skips_rows_with_no_budget_duration(table_name): + """A row with a past budget_reset_at but no budget_duration has no next + window to move to, so it stays due forever. Fetching it means re-reading and + re-zeroing it on every tick, and a full chunk of such rows makes the paged + scan report no progress and starve the whole phase. + """ + find_many = _run_reset_query(table_name) + + assert find_many.await_args.kwargs["where"]["NOT"] == {"budget_duration": None} + + +@pytest.mark.parametrize("table_name", ["user", "team", "budget", "key"]) +def test_get_data_reset_query_is_unlimited_when_no_limit_is_passed(table_name): + """Callers that pass no limit keep the old unbounded behaviour.""" + find_many = _run_reset_query(table_name) + + assert find_many.await_args.kwargs.get("take") is None @pytest.mark.parametrize("table_name", ["user", "team"]) @@ -1838,8 +1608,327 @@ def test_get_data_reset_query_selects_null_budget_reset_at(table_name): setattr(getattr(client.db, table_attr), "find_many", find_many) now = datetime.now(timezone.utc) - asyncio.run( - client.get_data(table_name=table_name, query_type="find_all", reset_at=now) - ) + asyncio.run(client.get_data(table_name=table_name, query_type="find_all", reset_at=now)) _asserts_null_reset_is_due(_extract_reset_where(find_many)) + + +def _key_row(token: str, budget_duration: Any = "30d"): + """A key that is already due for a reset, shaped like a get_data() row.""" + now = datetime.now(timezone.utc) + return type( + "LiteLLM_VerificationToken", + (), + { + "spend": 100.0, + "budget_duration": budget_duration, + "budget_reset_at": now - timedelta(hours=1), + "token": token, + }, + ) + + +def _user_row(user_id: str, budget_duration: Any = "30d"): + now = datetime.now(timezone.utc) + return type( + "LiteLLM_UserTable", + (), + { + "spend": 100.0, + "budget_duration": budget_duration, + "budget_reset_at": now - timedelta(hours=1), + "user_id": user_id, + }, + ) + + +def _team_row(team_id: str, budget_duration: Any = "30d"): + now = datetime.now(timezone.utc) + return type( + "LiteLLM_TeamTable", + (), + { + "spend": 100.0, + "budget_duration": budget_duration, + "budget_reset_at": now - timedelta(hours=1), + "team_id": team_id, + }, + ) + + +# --------------------------------------------------------------------------- +# Chunked batches +# --------------------------------------------------------------------------- + + +class ChunkedPrismaClient(MockPrismaClient): + """Replays a scripted sequence of get_data chunks per table. + + The last chunk repeats forever, so a phase that fails to terminate keeps + seeing rows rather than quietly running out of data. + """ + + def __init__(self, chunks_by_table: Dict[str, List[List[Any]]]): + super().__init__() + self._chunks_by_table = chunks_by_table + self.fetches_by_table: Dict[str, int] = {} + + async def get_data(self, table_name, query_type, **kwargs): + self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs}) + chunks = self._chunks_by_table.get(table_name) + if not chunks: + return [] + index = self.fetches_by_table.get(table_name, 0) + self.fetches_by_table[table_name] = index + 1 + return chunks[min(index, len(chunks) - 1)] + + +def _chunked_job(chunks_by_table): + client = ChunkedPrismaClient(chunks_by_table) + return client, ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + +def _fetch_limits(client, table_name): + return [call.get("limit") for call in client.get_data_calls if call["table_name"] == table_name] + + +def test_key_reset_walks_the_due_rows_one_chunk_at_a_time(monkeypatch): + """Each chunk is fetched under a LIMIT and committed on its own batch, so a + large backlog never becomes one giant transaction.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client, job = _chunked_job({"key": [[_key_row("k1"), _key_row("k2")], [_key_row("k3")]]}) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.fetches_by_table["key"] == 2 + assert _fetch_limits(client, "key") == [2, 2] + assert len(client.db.batchers) == 2 + assert all(batcher.committed for batcher in client.db.batchers) + assert [len(batcher.calls) for batcher in client.db.batchers] == [2, 1] + assert [w["where"]["token"] for w in _batch_writes(client, "key", op="update")] == ["k1", "k2", "k3"] + + +def test_key_reset_stops_after_a_chunk_shorter_than_the_batch_size(monkeypatch): + """Fewer rows than the limit means the backlog is drained, so no follow-up + query is worth issuing.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 5) + client, job = _chunked_job({"key": [[_key_row("k1")]]}) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.fetches_by_table["key"] == 1 + + +def test_key_reset_stops_when_a_full_chunk_advances_nothing(monkeypatch): + """A key with no budget_duration keeps its past budget_reset_at, so the very + same rows come back on the next fetch. Treating those writes as progress + would re-read that chunk until the iteration cap, every tick.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + stuck_chunk = [_key_row("k1", budget_duration=None), _key_row("k2", budget_duration=None)] + client, job = _chunked_job({"key": [stuck_chunk]}) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.fetches_by_table["key"] == 1 + + +def test_key_reset_stops_when_the_fetch_fails(monkeypatch): + """A phase whose query raises has made no progress; retrying it in a tight + loop would just hammer a struggling database.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client, job = _chunked_job({"key": [[_key_row("k1"), _key_row("k2")]]}) + + async def _boom(table_name, query_type, **kwargs): + client.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs}) + raise RuntimeError("db is down") + + client.get_data = _boom + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert len(client.get_data_calls) == 1 + + +def test_key_reset_is_capped_at_max_chunks_per_run(monkeypatch): + """Backstop against a phase that keeps making progress forever: the run ends + and the leftovers wait for the next tick.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 1) + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", 3) + client, job = _chunked_job({"key": [[_key_row("k1")]]}) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.fetches_by_table["key"] == 3 + + +@pytest.mark.parametrize( + "phase, table_name, row_factory", + [ + ("reset_budget_for_litellm_users", "user", lambda uid: _user_row(uid)), + ("reset_budget_for_litellm_teams", "team", lambda tid: _team_row(tid)), + ], + ids=["users", "teams"], +) +def test_user_and_team_resets_are_chunked_too(monkeypatch, phase, table_name, row_factory): + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client, job = _chunked_job({table_name: [[row_factory("a"), row_factory("b")], [row_factory("c")]]}) + + asyncio.run(getattr(job, phase)()) + + assert client.fetches_by_table[table_name] == 2 + assert _fetch_limits(client, table_name) == [2, 2] + assert len(client.db.batchers) == 2 + assert len(_batch_writes(client, table_name, op="update")) == 3 + + +def test_budget_table_reset_walks_chunks_until_it_runs_dry(monkeypatch): + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client, job = _chunked_job({"budget": [[_budget_row("b1"), _budget_row("b2")], [_budget_row("b3")]]}) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert client.fetches_by_table["budget"] == 2 + assert _fetch_limits(client, "budget") == [2, 2] + assert len(client.db.batchers) == 2 + assert all(batcher.committed for batcher in client.db.batchers) + assert [w["where"]["budget_id"] for w in _batch_writes(client, "budget", op="update_many")] == ["b1", "b2", "b3"] + + +def test_budget_table_reset_stops_when_a_full_chunk_advances_no_window(monkeypatch): + """A tier with no budget_duration has its linked spend zeroed but keeps its + past budget_reset_at, so it stays due. Counting those spend writes as + progress would re-read the same chunk until the cap.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + stuck_chunk = [_budget_row("b1", budget_duration=None), _budget_row("b2", budget_duration=None)] + client, job = _chunked_job({"budget": [stuck_chunk]}) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert client.fetches_by_table["budget"] == 1 + assert _batch_writes(client, "budget", op="update_many") == [] + + +def test_budget_table_reset_stops_when_the_cascade_fails(monkeypatch): + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client, job = _chunked_job({"budget": [[_budget_row("b1"), _budget_row("b2")]]}) + client.db = FailingCommitDB() + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert client.fetches_by_table["budget"] == 1 + + +# --------------------------------------------------------------------------- +# Progress means "no longer due", not "was written" +# --------------------------------------------------------------------------- + + +def test_key_reset_stops_when_the_new_reset_time_is_not_in_the_future(monkeypatch): + """A "0s" budget_duration resolves to the current time, so the row is written + and comes straight back on the next fetch. Treating a written row as progress + burns the whole per-run chunk cap on rows that never move. + """ + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + stuck_chunk = [_key_row("k1", budget_duration="0s"), _key_row("k2", budget_duration="0s")] + client, job = _chunked_job({"key": [stuck_chunk]}) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.fetches_by_table["key"] == 1 + assert len(_batch_writes(client, "key", op="update")) == 2 + + +def test_budget_table_reset_stops_when_the_new_window_is_not_in_the_future(monkeypatch): + """Same zero-length window on the budget tier: advancing it to now leaves it + due, so the cascade must not report progress.""" + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + stuck_chunk = [_budget_row("b1", budget_duration="0s"), _budget_row("b2", budget_duration="0s")] + client, job = _chunked_job({"budget": [stuck_chunk]}) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert client.fetches_by_table["budget"] == 1 + assert len(_batch_writes(client, "budget", op="update_many")) == 2 + + +class PoisonRow: + """A row the in-memory reset cannot write, like the DataError rows in #27730.""" + + token = "poison" + budget_duration = "30d" + budget_reset_at = None + + def __setattr__(self, name: str, value: Any) -> None: + raise RuntimeError("simulated failure resetting this row") + + +class RecordingServiceLogging: + def __init__(self): + self.success_calls: List[Dict[str, Any]] = [] + self.failure_calls: List[Dict[str, Any]] = [] + + async def async_service_success_hook(self, **kwargs): + self.success_calls.append(kwargs) + + async def async_service_failure_hook(self, **kwargs): + self.failure_calls.append(kwargs) + + +class RecordingProxyLogging: + def __init__(self): + self.service_logging_obj = RecordingServiceLogging() + + +def _run_and_drain_hooks(make_coro): + """The service hooks are fired as tasks; give them a turn before asserting.""" + + async def _run(): + await make_coro() + await asyncio.sleep(0.05) + + asyncio.run(_run()) + + +def test_key_reset_keeps_paging_when_some_rows_in_a_chunk_fail(monkeypatch): + """One row that cannot be reset must not cost the phase its remaining chunks: + the rows that did reset are committed and are real progress, and the failure + is reported instead of aborting the run. + """ + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client = ChunkedPrismaClient({"key": [[PoisonRow(), _key_row("k1")], [_key_row("k2")]]}) + logging_obj = RecordingProxyLogging() + job = ResetBudgetJob(proxy_logging_obj=logging_obj, prisma_client=client) + + _run_and_drain_hooks(job.reset_budget_for_litellm_keys) + + assert client.fetches_by_table["key"] == 2 + assert [w["where"]["token"] for w in _batch_writes(client, "key", op="update")] == ["k1", "k2"] + assert [call["call_type"] for call in logging_obj.service_logging_obj.failure_calls] == ["reset_budget_keys"] + assert set(logging_obj.service_logging_obj.failure_calls[0]["event_metadata"]) == { + "num_keys_found", + "keys_found", + } + assert [call["call_type"] for call in logging_obj.service_logging_obj.success_calls] == ["reset_budget_keys"] + + +@pytest.mark.parametrize( + "phase, table_name, row_factory, call_type", + [ + ("reset_budget_for_litellm_users", "user", _user_row, "reset_budget_users"), + ("reset_budget_for_litellm_teams", "team", _team_row, "reset_budget_teams"), + ], + ids=["users", "teams"], +) +def test_user_and_team_chunks_report_progress_despite_a_failed_row( + monkeypatch, phase, table_name, row_factory, call_type +): + monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) + client = ChunkedPrismaClient({table_name: [[PoisonRow(), row_factory("a")], [row_factory("b")]]}) + logging_obj = RecordingProxyLogging() + job = ResetBudgetJob(proxy_logging_obj=logging_obj, prisma_client=client) + + _run_and_drain_hooks(getattr(job, phase)) + + assert client.fetches_by_table[table_name] == 2 + assert len(_batch_writes(client, table_name, op="update")) == 2 + assert [call["call_type"] for call in logging_obj.service_logging_obj.failure_calls] == [call_type] diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 1c29c287c3a..076151fcd3b 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3,6 +3,7 @@ Unit Tests for the max parallel request limiter v3 for the proxy """ import asyncio +import logging import os import sys import time @@ -5100,3 +5101,456 @@ async def test_reserve_tpm_tokens_never_evaluates_the_requests_dimension(): f"reservation pass, got: {response}" ) assert [s["rate_limit_type"] for s in response["statuses"]] == ["tokens"] + + +STATIC_OUTPUT_FLOOR = 1024 +ONE_TOKEN_PROMPT = [{"role": "user", "content": "hello"}] +ONE_TOKEN_PROMPT_INPUT_ESTIMATE = 1 + + +async def _reserved_tokens_for( + handler, + local_cache, + user_api_key_dict, + data, + call_type="completion", +): + """Drive the pre-call hook and read back what landed on the :tokens counter.""" + tokens_key = handler.create_rate_limit_keys( + key="api_key", value=user_api_key_dict.api_key, rate_limit_type="tokens" + ) + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type=call_type, + ) + return int(await local_cache.async_get_cache(key=tokens_key) or 0) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "key_metadata, team_metadata, expected_output_estimate, tier", + [ + ( + { + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": 3001}, + "default_estimated_output_tokens": 2002, + }, + { + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503}, + "default_estimated_output_tokens": 777, + }, + 3001, + "key per-model wins over every other tier", + ), + ( + {"default_estimated_output_tokens": 2002}, + { + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503}, + "default_estimated_output_tokens": 777, + }, + 2002, + "key global wins over team config", + ), + ( + {"default_estimated_output_tokens_per_model": {"some-other-model": 9999}}, + { + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503}, + "default_estimated_output_tokens": 777, + }, + 1503, + "team per-model wins when the key has no applicable entry", + ), + ( + {}, + {"default_estimated_output_tokens": 777}, + 777, + "team global is the last configured tier", + ), + ({}, {}, STATIC_OUTPUT_FLOOR, "unconfigured falls back to the static floor"), + ( + {"unrelated": "value"}, + {"unrelated": "value"}, + STATIC_OUTPUT_FLOOR, + "unrelated metadata changes nothing", + ), + ( + {"default_estimated_output_tokens": "not-a-number"}, + {}, + STATIC_OUTPUT_FLOOR, + "malformed config falls back to the static floor instead of erroring", + ), + ( + {"default_estimated_output_tokens": 0}, + {}, + STATIC_OUTPUT_FLOOR, + "a non-positive estimate is rejected, not reserved", + ), + ], +) +async def test_estimated_output_tokens_resolution_precedence( + monkeypatch, key_metadata, team_metadata, expected_output_estimate, tier +): + """The no-max_tokens output reservation resolves per key / team / model. + + Every configured value here is distinct from the static 1024 floor and + from the input estimate, so the reserved amount identifies which tier the + resolver picked. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token(f"sk-estimate-{expected_output_estimate}-{tier}"), + tpm_limit=1_000_000, + metadata=key_metadata, + team_metadata=team_metadata, + ) + + reserved = await _reserved_tokens_for( + handler, + local_cache, + user_api_key_dict, + {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + ) + + assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + expected_output_estimate, tier + + +@pytest.mark.asyncio +async def test_request_max_tokens_outranks_configured_estimate(monkeypatch): + """An explicit request-level max_tokens stays the top of the precedence order.""" + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-estimate-explicit-max-tokens"), + tpm_limit=1_000_000, + metadata={"default_estimated_output_tokens": 2002}, + ) + + reserved = await _reserved_tokens_for( + handler, + local_cache, + user_api_key_dict, + {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT, "max_tokens": 42}, + ) + + assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 42 + + +@pytest.mark.asyncio +async def test_configured_estimate_does_not_apply_to_embeddings(monkeypatch): + """Embeddings generate no output, so a declared output estimate must not be reserved.""" + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-estimate-embeddings"), + tpm_limit=1_000_000, + metadata={"default_estimated_output_tokens": 2002}, + ) + + reserved = await _reserved_tokens_for( + handler, + local_cache, + user_api_key_dict, + {"model": "text-embedding-3-small", "input": "hello"}, + call_type="embeddings", + ) + + assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + + +@pytest.mark.asyncio +async def test_configured_estimate_applies_to_contentless_requests(monkeypatch): + """A declared estimate describes generation, so it holds even with no prompt body. + + Without config such a request reserves the 1-token floor only; the + declaration is what makes concurrent tool-call continuations countable. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + configured = UserAPIKeyAuth( + api_key=hash_token("sk-estimate-contentless-configured"), + tpm_limit=1_000_000, + metadata={"default_estimated_output_tokens": 2002}, + ) + unconfigured = UserAPIKeyAuth( + api_key=hash_token("sk-estimate-contentless-plain"), + tpm_limit=1_000_000, + ) + + assert ( + await _reserved_tokens_for( + handler, local_cache, configured, {"model": "gpt-4o-mini", "messages": []} + ) + == 2002 + ) + assert ( + await _reserved_tokens_for( + handler, local_cache, unconfigured, {"model": "gpt-4o-mini", "messages": []} + ) + == 1 + ) + + +@pytest.mark.asyncio +async def test_declared_estimate_never_tightens_the_small_tpm_clamp(monkeypatch): + """The small-TPM clamp can only be loosened by a declaration, never tightened. + + That clamp is the one place the proxy rewrites the caller's generation + budget, and it only fires below a 4096 TPM limit. A declaration above it + raises it, so the tenant is not truncated below what they said their + model emits; a declaration below it changes nothing, because an estimate + describes the typical response and must not become a hard cap that + truncates the tail. The reservation tracks whatever the clamp settles on, + so a small tenant can never generate more than was reserved. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + raised_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT} + raised_reserved = await _reserved_tokens_for( + handler, + local_cache, + UserAPIKeyAuth( + api_key=hash_token("sk-estimate-hard-cap-raised"), + tpm_limit=2000, + metadata={"default_estimated_output_tokens": 900}, + ), + raised_data, + ) + assert raised_data["max_tokens"] == 900 + assert raised_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 900 + + lowered_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT} + lowered_reserved = await _reserved_tokens_for( + handler, + local_cache, + UserAPIKeyAuth( + api_key=hash_token("sk-estimate-hard-cap-lowered"), + tpm_limit=2000, + metadata={"default_estimated_output_tokens": 120}, + ), + lowered_data, + ) + assert lowered_data["max_tokens"] == 500 + assert lowered_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 500 + + unconfigured_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT} + unconfigured_reserved = await _reserved_tokens_for( + handler, + local_cache, + UserAPIKeyAuth( + api_key=hash_token("sk-estimate-hard-cap-plain"), + tpm_limit=2000, + ), + unconfigured_data, + ) + assert unconfigured_data["max_tokens"] == 500 + assert unconfigured_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 500 + + +@pytest.mark.asyncio +async def test_one_malformed_estimate_field_does_not_discard_the_other(monkeypatch): + """Each declared field is validated on its own. + + A per-model map with a bad entry must not take a valid global estimate + down with it, and a bad global must not hide a valid per-model entry. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + broken_map = await _reserved_tokens_for( + handler, + local_cache, + UserAPIKeyAuth( + api_key=hash_token("sk-estimate-broken-map"), + tpm_limit=1_000_000, + metadata={ + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": "huge"}, + "default_estimated_output_tokens": 2002, + }, + ), + {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + ) + assert broken_map == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 2002 + + broken_global = await _reserved_tokens_for( + handler, + local_cache, + UserAPIKeyAuth( + api_key=hash_token("sk-estimate-broken-global"), + tpm_limit=1_000_000, + metadata={ + "default_estimated_output_tokens_per_model": {"gpt-4o-mini": 3001}, + "default_estimated_output_tokens": -5, + }, + ), + {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + ) + assert broken_global == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 3001 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("declared", [100_000, 5000]) +async def test_declared_estimate_over_the_tpm_budget_is_honored_and_explained(monkeypatch, caplog, declared): + """A declaration bigger than the budget must not be silently shrunk. + + Capping it against the TPM limit would re-admit exactly the traffic this + feature exists to hold back, so the request is refused instead and the + reservation is explained rather than leaving an unexplained 429 loop. + + ``declared == tpm_limit`` is the boundary case: the declaration alone + equals the limit, so only adding the input estimate tips the reservation + over. Comparing the declaration against the limit rather than the + reservation would refuse this request while saying nothing. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token(f"sk-estimate-over-budget-{declared}"), + tpm_limit=5000, + metadata={"default_estimated_output_tokens": declared}, + ) + + with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + explained = [ + record.getMessage() + for record in caplog.records + if "cannot be admitted even against an empty window" in record.getMessage() + ] + assert len(explained) == 1, f"expected exactly one explanation, got {explained}" + assert str(declared) in explained[0] + assert str(ONE_TOKEN_PROMPT_INPUT_ESTIMATE + declared) in explained[0] + assert "5000" in explained[0] + + +@pytest.mark.asyncio +async def test_a_key_that_declared_nothing_is_never_blamed_for_a_declaration(monkeypatch, caplog): + """A request can outgrow its budget on prompt size alone, with no declaration. + + The heuristic path reserves input plus the injected clamp, so a long + prompt against a small limit is refused without anyone having declared + anything. Blaming the declared field there would point an operator at a + setting they never set, to fix a 429 whose real cause is prompt size. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key=hash_token("sk-undeclared-long-prompt"), + tpm_limit=1000, + ), + cache=local_cache, + data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "x" * 3600}]}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + assert not [ + record for record in caplog.records if "cannot be admitted even against an empty window" in record.getMessage() + ] + + +@pytest.mark.asyncio +async def test_declared_estimate_inside_the_tpm_budget_is_not_explained(monkeypatch, caplog): + """The explanation is for requests that cannot fit, not for every request. + + Without this, a correctly configured key would emit one line per call. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): + await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key=hash_token("sk-estimate-within-budget"), + tpm_limit=5000, + metadata={"default_estimated_output_tokens": 1000}, + ), + cache=local_cache, + data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + call_type="completion", + ) + + assert not [ + record for record in caplog.records if "cannot be admitted even against an empty window" in record.getMessage() + ] + + +@pytest.mark.asyncio +async def test_configured_estimate_blocks_the_overrun_the_static_floor_admits(monkeypatch): + """Concurrent unbounded requests must stop at the declared budget. + + A key with tpm_limit=8000 whose model really emits ~3000 output tokens + admits 7 concurrent requests under the 1024 floor (7 * 1025 <= 8000), so + once they all report actual usage the window carries ~21000 tokens + against an 8000 limit. Declaring the real output size admits only the two + requests the budget actually covers. + """ + monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) + + async def admitted(metadata): + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token(f"sk-overrun-{metadata}"), + tpm_limit=8000, + metadata=metadata, + ) + accepted = 0 + for _ in range(10): + try: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}, + call_type="completion", + ) + except HTTPException: + break + accepted += 1 + return accepted + + assert await admitted({}) == 7 + assert await admitted({"default_estimated_output_tokens": 3000}) == 2 diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py index 167a06ed551..f3515e84d0d 100644 --- a/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py @@ -1,11 +1,18 @@ +import ast +import importlib.util +from pathlib import Path +from types import ModuleType from typing import Annotated +import fastapi.dependencies.utils as fastapi_dependency_utils +import pytest from fastapi import Depends, FastAPI, Header, Query, Request from fastapi.testclient import TestClient +import litellm.proxy.management_endpoints.management_v1.common as common_module from litellm.proxy.management_endpoints.management_v1.common import ( - ManagementProblem, PROBLEM_CONTENT_TYPE, + ManagementProblem, _declared_query_params, problem_response, reject_unknown_query_params, @@ -93,3 +100,57 @@ def test_declared_query_params_is_empty_when_the_route_has_no_dependant(): } ) assert _declared_query_params(request) == frozenset() + + +# fastapi removed these in 0.140.7, which `pyproject.toml` still allows via +# `fastapi>=0.136.3,<1.0`. Add a name here whenever a supported release drops one. +FASTAPI_NAMES_REMOVED_IN_0_140_7 = frozenset({"get_flat_dependant"}) + +MANAGEMENT_V1_PACKAGE = Path(str(common_module.__file__)).parent + + +def _public_names(module: ModuleType) -> frozenset[str]: + return frozenset(name for name in vars(module) if not name.startswith("_")) + + +def _fastapi_names_imported_by(source_file: Path) -> frozenset[str]: + tree = ast.parse(source_file.read_text()) + return frozenset( + alias.name + for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) and (node.module or "").startswith("fastapi") + for alias in node.names + ) + + +@pytest.mark.parametrize( + "source_file", sorted(MANAGEMENT_V1_PACKAGE.glob("*.py")), ids=lambda path: path.name +) +def test_no_module_imports_a_fastapi_name_removed_in_a_supported_release(source_file: Path): + """`pyproject.toml` allows fastapi up to <1.0, but CI only ever resolves 0.136.3. + + Every other test here passes just as well against a module importing a name + fastapi has since deleted, because the pinned fastapi still has it. On a user's + fastapi>=0.140.7 that import is an ImportError, and `proxy_server` imports this + package unguarded at module level, so it takes the whole proxy down rather than + just these routes. Globbing the package means a new module is covered on sight. + """ + assert not _fastapi_names_imported_by(source_file) & FASTAPI_NAMES_REMOVED_IN_0_140_7 + + +def test_common_still_imports_when_fastapi_has_dropped_those_names(monkeypatch: pytest.MonkeyPatch): + """The static check above cannot prove the module actually loads; this does. + + Behaviour cannot be asserted under the same simulation: on 0.136.3 + `get_flat_params` calls `get_flat_dependant` internally, so it raises NameError + once the name is gone. Loading is the part this pins. + """ + for name in FASTAPI_NAMES_REMOVED_IN_0_140_7: + monkeypatch.delattr(fastapi_dependency_utils, name, raising=False) + spec = importlib.util.spec_from_file_location( + "management_v1_common__simulated_fastapi", Path(str(common_module.__file__)) + ) + assert spec is not None and spec.loader is not None + reimported = importlib.util.module_from_spec(spec) + spec.loader.exec_module(reimported) + assert _public_names(reimported) == _public_names(common_module) diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index 3bdf9bafdc7..6a9e894feb5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -72,6 +72,39 @@ async def test_new_budget_success(client_and_mocks): mock_table.create.assert_awaited_once() +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +@pytest.mark.asyncio +async def test_new_budget_rejects_a_duration_that_never_advances( + client_and_mocks, bad_duration +): + """A zero-length window resets to "now", so the row is due again the moment + it is written and the reset job re-reads it on every tick forever.""" + client, _, mock_table = client_and_mocks + + resp = client.post( + "/budget/new", + json={"budget_id": "budget_bad", "max_budget": 10.0, "budget_duration": bad_duration}, + ) + + assert resp.status_code == 400, resp.text + assert "Invalid budget_duration" in resp.json()["detail"]["error"] + mock_table.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_budget_rejects_a_duration_that_never_advances(client_and_mocks): + client, _, mock_table = client_and_mocks + + resp = client.post( + "/budget/update", + json={"budget_id": "budget_456", "budget_duration": "0s"}, + ) + + assert resp.status_code == 400, resp.text + assert "Invalid budget_duration" in resp.json()["detail"]["error"] + mock_table.update.assert_not_awaited() + + @pytest.mark.asyncio async def test_new_budget_db_not_connected(client_and_mocks, monkeypatch): client, mock_prisma, mock_table = client_and_mocks diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index a388e7aaf09..ab9b4bc3922 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -7,6 +7,8 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path from litellm.proxy.management_endpoints.common_daily_activity import ( @@ -1142,6 +1144,11 @@ class TestEverySavingsDriverSurvivesTheReadPath: ) +@pytest.fixture +def ptu_cost_attribution_enabled(monkeypatch): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + + def _spend_record(api_key, *, model="gpt-4o-mini-ptu", spend=0.0, ptu_flat_cost=0.0): return SimpleNamespace( api_key=api_key, @@ -1167,13 +1174,13 @@ def _spend_record(api_key, *, model="gpt-4o-mini-ptu", spend=0.0, ptu_flat_cost= ) -def test_update_metrics_accumulates_ptu_flat_cost(): +def test_update_metrics_accumulates_ptu_flat_cost(ptu_cost_attribution_enabled): metrics = update_metrics(SpendMetrics(), _spend_record("real-key", spend=1.0, ptu_flat_cost=240.0)) assert metrics.flat_cost == 240.0 assert metrics.spend == 1.0 -def test_ptu_sentinel_excluded_from_key_breakdown_but_flat_cost_aggregates(): +def test_ptu_sentinel_excluded_from_key_breakdown_but_flat_cost_aggregates(ptu_cost_attribution_enabled): from litellm.constants import PTU_SENTINEL_API_KEY from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics @@ -1230,7 +1237,7 @@ def _grouping_row( ) -def test_grouping_sets_dispatcher_excludes_ptu_sentinel_from_key_breakdowns(): +def test_grouping_sets_dispatcher_excludes_ptu_sentinel_from_key_breakdowns(ptu_cost_attribution_enabled): """The GROUPING SETS path must mirror the per-row path: the flat-cost sentinel aggregates into the date/model/total metrics but never surfaces as an api_key.""" from litellm.constants import PTU_SENTINEL_API_KEY @@ -1267,7 +1274,7 @@ def test_grouping_sets_dispatcher_excludes_ptu_sentinel_from_key_breakdowns(): assert "real-key" in model_bucket.api_key_breakdown -def test_grouping_sets_dispatcher_populates_every_breakdown_level(): +def test_grouping_sets_dispatcher_populates_every_breakdown_level(ptu_cost_attribution_enabled): """Every GROUPING SETS level lands in its bucket, and the flat-cost sentinel is kept out of the model_group and provider api_key sub-breakdowns too.""" from litellm.constants import PTU_SENTINEL_API_KEY @@ -1359,7 +1366,7 @@ def test_grouping_sets_dispatcher_keeps_a_real_provider_row_that_shares_the_sent assert unknown.metrics.flat_cost == 0.0 -def test_update_breakdown_metrics_covers_mcp_endpoint_and_entity(): +def test_update_breakdown_metrics_covers_mcp_endpoint_and_entity(ptu_cost_attribution_enabled): """A full request record fans out into the mcp, endpoint, provider and entity breakdowns, while the flat-cost sentinel stays out of the entity api_key sub-map.""" from litellm.constants import PTU_SENTINEL_API_KEY @@ -1432,6 +1439,10 @@ class TestSentinelRowsDisplayTheirModelName: """A sentinel row keys on the deployment id so a rename cannot move it. The usage views render the breakdown key directly as a label, so the read path has to show the name.""" + @pytest.fixture(autouse=True) + def _enabled(self, ptu_cost_attribution_enabled): + """Flat cost is gated off by default, and these assert on the amounts.""" + @staticmethod def _breakdown(records): from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics @@ -1485,3 +1496,226 @@ class TestSentinelRowsDisplayTheirModelName: models = self._breakdown([self._sentinel(model_id="dep-1", model_group=None)]).models assert models["dep-1"].metrics.flat_cost == pytest.approx(480.0) + + +def _daily_team_row(api_key, *, spend=0.0, ptu_flat_cost=0.0): + """A LiteLLM_DailyTeamSpend row as the paginated read path receives it from find_many.""" + base: Final = _spend_record(api_key, spend=spend, ptu_flat_cost=ptu_flat_cost) + return SimpleNamespace(**{**base.__dict__, "date": "2026-07-01", "team_id": "team-1"}) + + +class TestPtuCostAttributionDisabled: + """With LITELLM_ENABLE_PTU_COST_ATTRIBUTION unset, both read paths report zero flat + cost, while the sentinel filtering that keeps ``__ptu_flat_cost__`` out of the + breakdowns keeps running. + + Filtering is deliberately not gated: an operator can enable the flag, accrue + sentinel rows, then disable it, and those rows stay in LiteLLM_DailyTeamSpend + forever. Gating the filter too would surface the sentinel as a bogus api_key and + mint a provider bucket for its empty provider. + """ + + @pytest.fixture(autouse=True) + def _flag_off(self, monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + + def test_paginated_path_reports_zero_flat_cost(self): + metrics = update_metrics(SpendMetrics(), _spend_record("real-key", spend=1.0, ptu_flat_cost=240.0)) + + assert metrics.flat_cost == 0.0 + assert metrics.spend == 1.0 + + def test_aggregated_path_reports_zero_flat_cost(self): + from litellm.proxy.management_endpoints.common_daily_activity import _GROUP_GRAND_TOTAL + + metrics = _record_to_spend_metrics(_grouping_row(_GROUP_GRAND_TOTAL, spend=5.0, ptu_flat_cost=240.0)) + + assert metrics.flat_cost == 0.0 + assert metrics.spend == 5.0 + + def test_aggregated_totals_and_buckets_report_zero_flat_cost(self): + from litellm.constants import PTU_SENTINEL_API_KEY + from litellm.proxy.management_endpoints.common_daily_activity import ( + _GROUP_DATE_API_KEY, + _GROUP_DATE_MODEL, + _GROUP_GRAND_TOTAL, + _aggregate_grouping_sets_records_sync, + ) + + records = [ + _grouping_row(_GROUP_DATE_API_KEY, api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), + _grouping_row(_GROUP_DATE_MODEL, model="gpt-4o-mini-ptu", spend=5.0, ptu_flat_cost=240.0), + _grouping_row(_GROUP_GRAND_TOTAL, spend=5.0, ptu_flat_cost=240.0), + ] + + aggregated = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={}) + + assert aggregated["totals"].flat_cost == 0.0 + assert aggregated["totals"].spend == 5.0 + assert aggregated["results"][0].breakdown.models["gpt-4o-mini-ptu"].metrics.flat_cost == 0.0 + + def test_sentinel_still_excluded_from_the_api_key_breakdown(self): + from litellm.constants import PTU_SENTINEL_API_KEY + from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics + from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics + + breakdown = BreakdownMetrics() + update_breakdown_metrics(breakdown, _spend_record("real-key", spend=5.0), {}, {}, {}) + update_breakdown_metrics( + breakdown, _spend_record(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), {}, {}, {}, entity_id_field="team_id" + ) + + assert PTU_SENTINEL_API_KEY not in breakdown.api_keys + assert PTU_SENTINEL_API_KEY not in breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown + assert "real-key" in breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown + + def test_sentinel_still_excluded_from_the_provider_breakdown(self): + from litellm.constants import PTU_SENTINEL_API_KEY + from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics + from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics + + breakdown = BreakdownMetrics() + update_breakdown_metrics(breakdown, _spend_record(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), {}, {}, {}) + + assert breakdown.providers == {} + + def test_grouping_sets_sentinel_still_excluded_from_breakdowns(self): + from litellm.constants import PTU_SENTINEL_API_KEY + from litellm.proxy.management_endpoints.common_daily_activity import ( + _GROUP_DATE_API_KEY, + _GROUP_DATE_MODEL, + _GROUP_DATE_MODEL_API_KEY, + _GROUP_DATE_PROVIDER, + _aggregate_grouping_sets_records_sync, + ) + + records = [ + _grouping_row(_GROUP_DATE_API_KEY, api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), + _grouping_row(_GROUP_DATE_MODEL, model="gpt-4o-mini-ptu", spend=5.0, ptu_flat_cost=240.0), + _grouping_row( + _GROUP_DATE_MODEL_API_KEY, model="gpt-4o-mini-ptu", api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0 + ), + _grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="", ptu_flat_cost=240.0), + ] + + day = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})["results"][0] + + assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys + assert PTU_SENTINEL_API_KEY not in day.breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown + assert sum(bucket.metrics.flat_cost for bucket in day.breakdown.providers.values()) == 0.0 + + @pytest.mark.asyncio + async def test_team_daily_activity_endpoint_reports_zero_flat_cost(self): + """/team/daily/activity reads rows with find_many rather than the aggregated SQL, so + forcing the SQL select to a constant zero would leave this path reporting flat cost.""" + from litellm.constants import PTU_SENTINEL_API_KEY + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=2) + mock_table.find_many = AsyncMock( + return_value=[ + _daily_team_row("real-key", spend=5.0), + _daily_team_row(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), + ] + ) + mock_prisma.db.litellm_verificationtoken = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_dailyteamspend = mock_table + + result = await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyteamspend", + entity_id_field="team_id", + entity_id="team-1", + entity_metadata_field=None, + start_date="2026-07-01", + end_date="2026-07-01", + model=None, + api_key=None, + page=1, + page_size=50, + ) + + assert result.metadata.total_flat_cost == 0.0 + assert result.metadata.total_spend == 5.0 + assert PTU_SENTINEL_API_KEY not in result.results[0].breakdown.api_keys + + @pytest.mark.asyncio + async def test_team_daily_activity_endpoint_reports_flat_cost_once_enabled(self, monkeypatch): + from litellm.constants import PTU_SENTINEL_API_KEY + + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=2) + mock_table.find_many = AsyncMock( + return_value=[ + _daily_team_row("real-key", spend=5.0), + _daily_team_row(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), + ] + ) + mock_prisma.db.litellm_verificationtoken = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_dailyteamspend = mock_table + + result = await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyteamspend", + entity_id_field="team_id", + entity_id="team-1", + entity_metadata_field=None, + start_date="2026-07-01", + end_date="2026-07-01", + model=None, + api_key=None, + page=1, + page_size=50, + ) + + assert result.metadata.total_flat_cost == 240.0 + assert PTU_SENTINEL_API_KEY not in result.results[0].breakdown.api_keys + + +class TestFlagIsNotReadOnTheHotPath: + """update_metrics runs once per accumulation and a record fans out across roughly a + dozen breakdowns, so a flag that reads through the secret manager must not be consulted + for rows that carry no flat cost at all.""" + + @staticmethod + def _count_flag_reads(records): + import litellm.proxy.management_endpoints.common_daily_activity as cda + from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics + + reads = [] + real = cda.is_ptu_cost_attribution_enabled + + def counted(): + reads.append(1) + return real() + + cda.is_ptu_cost_attribution_enabled = counted + try: + breakdown = BreakdownMetrics() + for record in records: + cda.update_breakdown_metrics(breakdown, record, {}, {}, {}) + finally: + cda.is_ptu_cost_attribution_enabled = real + return len(reads) + + def test_a_request_row_never_reads_the_flag(self): + reads = self._count_flag_reads([_spend_record("real-key", spend=5.0, ptu_flat_cost=0.0)]) + assert reads == 0, f"{reads} secret-manager lookups for a row with no flat cost" + + def test_a_page_of_request_rows_never_reads_the_flag(self): + rows = [_spend_record(f"key-{i}", spend=1.0, ptu_flat_cost=0.0) for i in range(50)] + assert self._count_flag_reads(rows) == 0 + + def test_a_sentinel_row_still_consults_the_flag(self): + from litellm.constants import PTU_SENTINEL_API_KEY + + reads = self._count_flag_reads([_spend_record(PTU_SENTINEL_API_KEY, spend=0.0, ptu_flat_cost=240.0)]) + assert reads > 0 diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index 81840745d0e..7dfd99dfa53 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -628,6 +628,58 @@ class TestValidateFiniteSpendErrorDetail: } +class TestValidateBudgetDuration: + """`validate_budget_duration` keeps durations that never advance out of the + database. + + A duration of "0s" resolves to a reset time of now, so the row is due again + the instant it is written. The reset job re-reads such rows on every tick + and, once one tenant owns enough of them, they fill each batch and starve + every other tenant's reset. + """ + + def test_none_is_allowed(self): + from litellm.proxy.management_endpoints.common_utils import ( + validate_budget_duration, + ) + + assert validate_budget_duration(None) is None + + @pytest.mark.parametrize("duration", ["30s", "5m", "1h", "1d", "7d", "30d", "1mo"]) + def test_positive_durations_are_allowed(self, duration): + from litellm.proxy.management_endpoints.common_utils import ( + validate_budget_duration, + ) + + assert validate_budget_duration(duration) is None + + @pytest.mark.parametrize("duration", ["0s", "0m", "0h", "0d", "-5m", "abc", ""]) + def test_non_advancing_durations_are_rejected(self, duration): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + validate_budget_duration, + ) + + with pytest.raises(HTTPException) as exc_info: + validate_budget_duration(duration) + assert exc_info.value.status_code == 400 + + def test_rejection_detail_is_exact(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + validate_budget_duration, + ) + + with pytest.raises(HTTPException) as exc_info: + validate_budget_duration("0s") + + assert exc_info.value.detail == { + "error": "Invalid budget_duration '0s'. Use a format like '1h', '24h', '7d', or '30d'." + } + + class TestRequireCallerUserIdErrorDetail: """The 403 for a service-account key must carry the exact error body.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 0af5ad6cd9b..5efed8de325 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -749,6 +749,40 @@ def test_char_new_body(mock_prisma_client, mock_user_api_key_auth): assert response.json() == _EXPECTED_CUSTOMER +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +def test_customer_new_rejects_a_duration_that_never_advances( + mock_prisma_client, mock_user_api_key_auth, bad_duration +): + """A zero-length window resets to "now", leaving the customer's budget row + permanently due for the reset job to re-read every tick.""" + mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + + response = client.post( + "/customer/new", + json={"user_id": "c1", "max_budget": 10.0, "budget_duration": bad_duration}, + headers={"Authorization": "Bearer k"}, + ) + + assert response.status_code == 400, response.text + assert "Invalid budget_duration" in response.text + mock_prisma_client.db.litellm_endusertable.create.assert_not_awaited() + + +def test_customer_new_accepts_a_normal_duration(mock_prisma_client, mock_user_api_key_auth): + mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.db.litellm_budgettable.create = AsyncMock( + return_value=_row({"budget_id": "b1", "max_budget": 10.0}) + ) + + response = client.post( + "/customer/new", + json={"user_id": "c1", "max_budget": 10.0, "budget_duration": "30d"}, + headers={"Authorization": "Bearer k"}, + ) + + assert response.status_code == 200, response.text + + def test_char_update_body(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=_row({"user_id": "c1", "blocked": False}) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 056c2d3657a..cd5a5d42b09 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -788,6 +788,68 @@ def test_update_internal_user_params_reset_spend_and_max_budget(): assert "budget_duration" not in non_default_values # Should not add default values +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +def test_update_internal_user_params_rejects_a_duration_that_never_advances(bad_duration): + """A zero-length window resets to "now", so the user row is due again the + moment it is written and the reset job re-reads it on every tick. Enough of + them fill each batch and starve other tenants' resets. + """ + from fastapi import HTTPException + + from litellm.proxy._types import UpdateUserRequest + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_internal_user_params, + ) + + data = UpdateUserRequest(user_id="test_user_id", budget_duration=bad_duration) + + with pytest.raises(HTTPException) as exc_info: + _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data) + + assert exc_info.value.status_code == 400 + assert "Invalid budget_duration" in str(exc_info.value.detail) + + +def test_update_internal_user_params_accepts_a_normal_duration(): + from litellm.proxy._types import UpdateUserRequest + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_internal_user_params, + ) + + data = UpdateUserRequest(user_id="test_user_id", budget_duration="30d") + + non_default_values = _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data) + + assert non_default_values["budget_duration"] == "30d" + assert non_default_values["budget_reset_at"] is not None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +async def test_new_user_rejects_a_duration_that_never_advances(mocker, bad_duration): + """/user/new must reject the same never-advancing durations /user/update does.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.management_endpoints.internal_user_endpoints import new_user + + mocker.patch("litellm.proxy.proxy_server.prisma_client", MagicMock()) + duplicate_check = mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + new=AsyncMock(), + ) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with pytest.raises(ProxyException) as exc_info: + await new_user( + data=NewUserRequest(budget_duration=bad_duration), + user_api_key_dict=admin, + ) + + assert str(exc_info.value.code) == "400" + assert "Invalid budget_duration" in str(exc_info.value.message) + duplicate_check.assert_not_awaited() + + @pytest.mark.asyncio async def test_new_user_license_over_limit(mocker): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index e8709f3af34..8f151ed882c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -2527,6 +2527,72 @@ def _setup_update_key_mocks(monkeypatch, mock_prisma_client): monkeypatch.setattr("litellm.store_audit_logs", False) +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +async def test_update_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration): + """A zero-length window resets to "now", so the key row is due again the + moment it is written. The reset job re-reads such rows on every tick, and a + tenant with enough of them fills each batch and starves other tenants. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken(token=hashed_token, user_id="test-user") + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + mock_prisma_client.update_data = AsyncMock() + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key=hashed_token, budget_duration=bad_duration), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "400" + assert "Invalid budget_duration" in str(exc_info.value.message) + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +async def test_generate_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration): + """/key/generate must reject the same never-advancing durations /key/update does.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_fn, + ) + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new=AsyncMock(), + ) as mock_generate: + with pytest.raises(ProxyException) as exc_info: + await generate_key_fn( + data=GenerateKeyRequest(budget_duration=bad_duration), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234" + ), + ) + + assert str(exc_info.value.code) == "400" + assert "Invalid budget_duration" in str(exc_info.value.message) + mock_generate.assert_not_awaited() + + @pytest.mark.asyncio async def test_update_key_by_alias_only(monkeypatch): """ @@ -8081,7 +8147,7 @@ async def test_key_with_budget_id_does_not_store_budget_duration(): budget_duration, the key does NOT get budget_duration stored on it. Keys with budget_id follow their linked budget tier's reset schedule; - reset_budget_for_keys_linked_to_budgets() resets them when the tier resets. + reset_budget_for_litellm_budget_table() resets them when the tier resets. This avoids duplicating budget_duration on keys so tier updates apply automatically to all linked keys. """ @@ -15442,3 +15508,312 @@ async def test_migrate_encryption_endpoint_rejects_proxy_admin_viewer(): assert exc_info.value.status_code == 403 mock_migrate.assert_not_awaited() + + +_ESTIMATE = "default_estimated_output_tokens" +_ESTIMATE_PER_MODEL = "default_estimated_output_tokens_per_model" + + +@pytest.mark.parametrize( + "label, request_body, existing_metadata, allowed", + [ + ("nothing declared", {}, None, True), + ("declared top-level on a key with none stored", {_ESTIMATE: 1}, None, False), + ("declared inside metadata on a key with none stored", {"metadata": {_ESTIMATE: 1}}, None, False), + ( + "per-model map declared inside metadata", + {"metadata": {_ESTIMATE_PER_MODEL: {"gpt-4": 1}}}, + None, + False, + ), + ("unrelated edit, metadata omitted", {"models": ["gpt-4"]}, {_ESTIMATE: 2000}, True), + ("stored value resent unchanged", {_ESTIMATE: 2000}, {_ESTIMATE: 2000}, True), + ("stored value lowered", {_ESTIMATE: 1}, {_ESTIMATE: 2000}, False), + ("stored value raised", {_ESTIMATE: 9000}, {_ESTIMATE: 2000}, False), + ( + "stored value cleared by sending a metadata blob without it", + {"metadata": {"other": "keep"}}, + {_ESTIMATE: 2000, "other": "keep"}, + False, + ), + ( + "stored value resent inside the metadata blob", + {"metadata": {_ESTIMATE: 2000, "other": "keep"}}, + {_ESTIMATE: 2000, "other": "keep"}, + True, + ), + ( + "per-model map resent unchanged", + {_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + {_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + True, + ), + ( + "one model in the per-model map lowered", + {_ESTIMATE_PER_MODEL: {"gpt-4": 1}}, + {_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + False, + ), + ], +) +def test_output_token_estimate_admin_gate_matrix(label, request_body, existing_metadata, allowed): + """A non-admin may only leave a key's stored output-token estimate exactly as it is. + + The estimate decides what the TPM limiter reserves for a request that omits + max_tokens, so lowering, raising or clearing it moves a reservation charged + against team and organization windows the key holder does not own. Key + metadata is writable by the key holder, and the declaration can be written + either as a dedicated top-level field or nested in the metadata blob, so + both routes are gated. Resending the stored value is what the edit form + produces on every save and has to stay allowed. + """ + from litellm.proxy.auth.auth_utils import ( + enforce_output_token_estimates_are_admin_only, + ) + + def _call(caller): + enforce_output_token_estimates_are_admin_only( + data=UpdateKeyRequest(key="sk-1", **request_body), + existing_metadata=existing_metadata, + user_api_key_dict=caller, + entity="key", + ) + + non_admin = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-non-admin", + user_id="alice", + ) + if allowed: + _call(non_admin) + else: + 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) + + _call( + UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + ) + + +@pytest.mark.asyncio +async def test_generate_key_output_token_estimate_rejected_for_non_admin(): + """The /key/update gate does not cover generate, so without its own check a + non-admin could self-mint a key that reserves one output token per + unbounded request and overrun the TPM window it is charged against.""" + with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()): + with pytest.raises(HTTPException) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest(default_estimated_output_tokens=1, tpm_limit=100000), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + litellm_changed_by=None, + team_table=None, + ) + assert int(getattr(exc.value, "status_code", 0)) == 403 + assert "Only proxy admins can set" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_generate_key_output_token_estimate_in_metadata_rejected_for_non_admin(): + """Writing the declaration into the raw metadata blob lands in the same + stored field, so gating only the dedicated top-level field leaves the + bypass wide open.""" + with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()): + with pytest.raises(HTTPException) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest(metadata={"default_estimated_output_tokens": 1}), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + litellm_changed_by=None, + team_table=None, + ) + assert int(getattr(exc.value, "status_code", 0)) == 403 + + +@pytest.mark.asyncio +async def test_generate_key_output_token_estimate_allowed_for_admin(): + """A proxy admin declaring the estimate must reach key creation.""" + with ( + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", False), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn" + ) as mock_generate_key, + ): + mock_generate_key.return_value = { + "key": "sk-test-key", + "expires": None, + "user_id": "admin", + "team_id": None, + } + await _common_key_generation_helper( + data=GenerateKeyRequest(default_estimated_output_tokens=200), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"), + litellm_changed_by=None, + team_table=None, + ) + assert mock_generate_key.called + + +def _estimate_key_row(token: str, metadata: dict): + existing_key = MagicMock() + existing_key.token = token + existing_key.user_id = "internal_user" + existing_key.created_by = "internal_user" + existing_key.team_id = None + existing_key.project_id = None + existing_key.max_budget = 10.0 + existing_key.key_alias = None + existing_key.models = [] + existing_key.metadata = metadata + existing_key.model_dump.return_value = { + "token": token, + "user_id": "internal_user", + "team_id": None, + "max_budget": 10.0, + } + return existing_key + + +def _wire_update_key_fn(monkeypatch, existing_key): + mock_prisma_client = AsyncMock() + updated_key = MagicMock() + updated_key.token = existing_key.token + updated_key.key_alias = "my-alias" + + mock_prisma_client.get_data = AsyncMock(return_value=existing_key) + mock_prisma_client.update_data = AsyncMock(return_value=updated_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", lambda token: existing_key.token) + + async def _noop(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + _noop, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", + _noop, + ) + + +@pytest.mark.asyncio +async def test_update_key_output_token_estimate_lowered_rejected_for_non_admin(monkeypatch): + """End-to-end wiring: a key's owner reaches /key/update without any admin + check because metadata is a non-budget field, so the gate has to fire + inside the update path itself rather than only in a helper.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + token = "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + _wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_ESTIMATE: 4000})) + + mock_request = MagicMock() + mock_request.query_params = {} + + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=token, default_estimated_output_tokens=1), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + ) + + assert str(exc.value.code) == "403" + assert "Only proxy admins can set" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_update_key_output_token_estimate_unchanged_allows_non_admin_edit(monkeypatch): + """The edit form resends every field it renders, so gating on presence + would 403 a key owner renaming their own key.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + token = "b1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + _wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_ESTIMATE: 4000})) + + mock_request = MagicMock() + mock_request.query_params = {} + + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=token, key_alias="my-alias", default_estimated_output_tokens=4000), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_regenerate_key_output_token_estimate_lowered_rejected_for_non_admin(): + """/key/regenerate is a third write path into the same stored metadata. + + can_modify_verification_token lets a key's own holder regenerate it, and + the request body runs through prepare_key_update_data exactly as an update + does, so gating only generate and update leaves the declaration writable. + """ + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + + token = "c1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + key_in_db = LiteLLM_VerificationToken( + token=token, + user_id="internal_user", + metadata={_ESTIMATE: 4000}, + ) + + with pytest.raises(HTTPException) as exc: + await _execute_virtual_key_regeneration( + prisma_client=AsyncMock(), + key_in_db=key_in_db, + hashed_api_key=token, + key="sk-original", + data=RegenerateKeyRequest(key="sk-original", default_estimated_output_tokens=1), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert exc.value.status_code == 403 + assert "Only proxy admins can set" in str(exc.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py index e8131854acf..30fe78d93c7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py @@ -1,34 +1,57 @@ -import datetime -import json - """Tests for PTU config on the model deployment (v1 model-settings design).""" -from unittest.mock import AsyncMock, MagicMock +import datetime +import json +from contextlib import ExitStack +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException +from litellm.proxy._types import LiteLLM_ProxyModelTable, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.model_management_endpoints import ( _merged_ptu_model_info, + _raise_if_ptu_cost_attribution_disabled, _validate_ptu_model_info, + add_new_model, + update_db_model, ) +from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo, updateDeployment def test_model_info_accepts_valid_ptu_fields(): - info = ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=2.0) + info = ModelInfo( + id="x", + team_id="t", + ptu_count=5, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ) assert info.ptu_count == 5 assert info.cost_per_ptu_per_hour == 2.0 def test_model_info_rejects_non_positive_count(): with pytest.raises(ValueError): - ModelInfo(id="x", team_id="t", ptu_count=0, cost_per_ptu_per_hour=2.0) + ModelInfo( + id="x", + team_id="t", + ptu_count=0, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ) def test_model_info_rejects_negative_rate(): with pytest.raises(ValueError): - ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=-1.0) + ModelInfo( + id="x", + team_id="t", + ptu_count=5, + cost_per_ptu_per_hour=-1.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ) def test_model_info_rejects_a_count_beyond_the_cap(): @@ -223,6 +246,11 @@ class TestPartialPtuEditsUseTheMergedView: """A PTU invariant holds over the deployment as it will exist, not over whichever subset of fields a caller sent. Validating the patch alone rejected an ordinary edit.""" + @pytest.fixture(autouse=True) + def _enabled(self, monkeypatch): + """PTU writes are gated off by default; these are about the validator, not the gate.""" + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + @staticmethod def _configured(): return Deployment( @@ -310,6 +338,11 @@ class TestTeamModelUpdateValidatesBeforeWriting: """Drives the endpoint path itself, not the helpers. The validator sits above the team ACL write, which autocommits, so what it validates has to be right at that call site.""" + @pytest.fixture(autouse=True) + def _enabled(self, monkeypatch): + """PTU writes are gated off by default; these are about the validator, not the gate.""" + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + @staticmethod async def _run(db_model, patch_data, monkeypatch, touched=None): import litellm.proxy.management_endpoints.model_management_endpoints as mme @@ -352,6 +385,33 @@ class TestTeamModelUpdateValidatesBeforeWriting: assert "ptu_effective_from is required" in exc.value.detail + @pytest.mark.asyncio + async def test_the_gate_refuses_before_the_team_write(self, monkeypatch): + """The gate lived inside update_db_model, which runs after the team ACL write, so a + rejected edit still moved the model between teams.""" + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + db_model = Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params(model="openai/gpt-4o"), + model_info=ModelInfo(id="dep-0", team_id="team-A"), + ) + patch = updateDeployment( + model_info=ModelInfo( + id="dep-0", + team_id="team-B", + ptu_count=15, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2026, 8, 1, tzinfo=datetime.timezone.utc), + ) + ) + touched = [] + + with pytest.raises(HTTPException) as exc: + await self._run(db_model, patch, monkeypatch, touched) + + assert PTU_COST_ATTRIBUTION_ENV_VAR in exc.value.detail + assert touched == [] + @pytest.mark.asyncio async def test_clearing_half_the_pair_is_refused_before_the_team_write(self, monkeypatch): """The write drops the nulled field, so validating against the stored one let a @@ -380,3 +440,270 @@ class TestTeamModelUpdateValidatesBeforeWriting: stored = json.loads(result["model_info"]) assert "ptu_count" not in stored assert "cost_per_ptu_per_hour" not in stored + + +class TestPtuCostAttributionGate: + """PTU config is only writable once an operator sets LITELLM_ENABLE_PTU_COST_ATTRIBUTION. + + The fields are rejected rather than dropped: a silent accept-and-drop would let a + caller believe a flat cost was configured while the rollup that prices it is not + even scheduled. + """ + + @pytest.fixture(autouse=True) + def _flag_off(self, monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + + @pytest.fixture + def flag_on(self, monkeypatch): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + + @pytest.mark.parametrize( + "model_info", + [ + {"team_id": "t", "ptu_count": 5, "cost_per_ptu_per_hour": 2.0}, + {"ptu_count": 5}, + {"cost_per_ptu_per_hour": 2.0}, + {"ptu_effective_from": "2026-08-01T00:00:00Z"}, + {"ptu_effective_to": "2026-08-02T00:00:00Z"}, + ], + ) + def test_rejects_any_ptu_field_while_disabled(self, model_info): + with pytest.raises(HTTPException) as exc: + _raise_if_ptu_cost_attribution_disabled(model_info) + assert exc.value.status_code == 400 + assert PTU_COST_ATTRIBUTION_ENV_VAR in exc.value.detail + + def test_names_every_offending_field(self): + with pytest.raises(HTTPException) as exc: + _raise_if_ptu_cost_attribution_disabled({"ptu_count": 5, "cost_per_ptu_per_hour": 2.0}) + assert "ptu_count" in exc.value.detail + assert "cost_per_ptu_per_hour" in exc.value.detail + + def test_allows_a_request_without_ptu_fields_while_disabled(self): + _raise_if_ptu_cost_attribution_disabled({"team_id": "t", "access_groups": ["a"]}) + + def test_allows_every_ptu_field_once_enabled(self, flag_on): + _raise_if_ptu_cost_attribution_disabled( + { + "team_id": "t", + "ptu_count": 5, + "cost_per_ptu_per_hour": 2.0, + "ptu_effective_from": "2026-08-01T00:00:00Z", + "ptu_effective_to": "2026-08-02T00:00:00Z", + } + ) + + +def _deployment_without_ptu() -> Deployment: + return Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params(model="openai/gpt-4o"), + model_info=ModelInfo(id="dep-0", team_id="t"), + ) + + +def _deployment_with_stored_ptu() -> Deployment: + return Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params(model="openai/gpt-4o"), + model_info=ModelInfo( + id="dep-0", + team_id="t", + ptu_count=15, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ), + ) + + +class TestUpdateDbModelPtuGate: + @pytest.fixture(autouse=True) + def _flag_off(self, monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + + def test_patch_carrying_ptu_config_is_rejected(self): + with pytest.raises(HTTPException) as exc: + update_db_model( + db_model=_deployment_without_ptu(), + updated_patch=updateDeployment(model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=15)), + ) + assert exc.value.status_code == 400 + + def test_patch_that_touches_nothing_ptu_still_succeeds(self): + result = update_db_model( + db_model=_deployment_without_ptu(), + updated_patch=updateDeployment(model_info=ModelInfo(id="dep-0", access_groups=["a"])), + ) + assert json.loads(result["model_info"])["access_groups"] == ["a"] + + def test_unrelated_patch_of_a_model_that_stores_ptu_config_is_not_blocked(self): + """A deployment configured during an earlier opt-in stays editable: the gate reads the + incoming patch, not the merged deployment, so the stored config is left in place.""" + result = update_db_model( + db_model=_deployment_with_stored_ptu(), + updated_patch=updateDeployment(model_name="gpt-4o-renamed"), + ) + assert result["model_name"] == "gpt-4o-renamed" + + def test_explicit_nulls_do_not_erase_stored_ptu_config_while_disabled(self): + """A client round-tripping a model_info blob sends the PTU keys as nulls. While the + feature is disabled those nulls must not reach the clear loop: disabling pauses PTU, + it does not silently discard a billing configuration the operator set up earlier.""" + result = update_db_model( + db_model=_deployment_with_stored_ptu(), + updated_patch=updateDeployment( + model_info=ModelInfo(id="dep-0", ptu_count=None, cost_per_ptu_per_hour=None) + ), + ) + stored = json.loads(result["model_info"]) + assert stored["ptu_count"] == 15 + assert stored["cost_per_ptu_per_hour"] == 2.0 + + def test_the_merged_view_agrees_with_the_write_while_disabled(self): + """The validator sees what the write will store. If the merged view honoured a null the + clear loop ignores, a round-tripped blob would 400 on a half-set pair that never forms.""" + merged = _merged_ptu_model_info( + db_model=_deployment_with_stored_ptu(), + patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=None)), + ) + assert merged["ptu_count"] == 15 + _validate_ptu_model_info(merged) + + def test_explicit_nulls_still_clear_once_enabled(self, monkeypatch): + """Clearing remains available to an operator who opted in, which is how PTU config is + removed from a deployment.""" + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + result = update_db_model( + db_model=_deployment_with_stored_ptu(), + updated_patch=updateDeployment( + model_info=ModelInfo(id="dep-0", ptu_count=None, cost_per_ptu_per_hour=None) + ), + ) + stored = json.loads(result["model_info"]) + assert "ptu_count" not in stored + assert "cost_per_ptu_per_hour" not in stored + + def test_patch_carrying_ptu_config_is_accepted_once_enabled(self, monkeypatch): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + result = update_db_model( + db_model=_deployment_without_ptu(), + updated_patch=updateDeployment( + model_info=ModelInfo( + id="dep-0", + team_id="t", + ptu_count=15, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ) + ), + ) + stored = json.loads(result["model_info"]) + assert stored["ptu_count"] == 15 + assert stored["cost_per_ptu_per_hour"] == 2.0 + + +class TestAddNewModelPtuGate: + @pytest.fixture(autouse=True) + def _flag_off(self, monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + + @staticmethod + def _patched_proxy(model_id: str): + """Patch everything /model/new touches except the PTU gate, and hand back the DB writers.""" + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name="ptu-model", + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={"id": model_id}, + created_by="test-admin", + updated_by="test-admin", + ) + add_model_to_db = AsyncMock(return_value=db_row) + add_team_model_to_db = AsyncMock(return_value=db_row) + + mock_proxy_config = MagicMock() + mock_proxy_config.add_deployment = AsyncMock(return_value=None) + + mock_router = MagicMock() + mock_router.get_model_ids.return_value = [model_id] + + proxy_server = "litellm.proxy.proxy_server" + endpoints = "litellm.proxy.management_endpoints.model_management_endpoints" + return (add_model_to_db, add_team_model_to_db), [ + patch(f"{proxy_server}.prisma_client", MagicMock()), + patch(f"{proxy_server}.store_model_in_db", True), + patch(f"{proxy_server}.proxy_config", mock_proxy_config), + patch(f"{proxy_server}.proxy_logging_obj", MagicMock()), + patch(f"{proxy_server}.general_settings", {}), + patch(f"{proxy_server}.premium_user", True), + patch(f"{proxy_server}.llm_router", mock_router), + patch( + f"{endpoints}.ModelManagementAuthChecks.can_user_make_model_call", + AsyncMock(return_value=True), + ), + patch(f"{endpoints}._add_model_to_db", add_model_to_db), + patch(f"{endpoints}._add_team_model_to_db", add_team_model_to_db), + ] + + @staticmethod + def _ptu_deployment(model_id: str) -> Deployment: + return Deployment( + model_name="ptu-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"), + model_info=ModelInfo( + id=model_id, + team_id="team-1", + ptu_count=15, + cost_per_ptu_per_hour=2.0, + ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc), + ), + ) + + @pytest.mark.asyncio + async def test_model_new_rejects_ptu_config_while_disabled(self): + (add_model_to_db, add_team_model_to_db), patches = self._patched_proxy("ptu-gate-model") + admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ExitStack() as stack: + for active_patch in patches: + stack.enter_context(active_patch) + with pytest.raises(Exception) as exc: + await add_new_model(model_params=self._ptu_deployment("ptu-gate-model"), user_api_key_dict=admin) + + assert PTU_COST_ATTRIBUTION_ENV_VAR in str(exc.value) + add_model_to_db.assert_not_called() + add_team_model_to_db.assert_not_called() + + @pytest.mark.asyncio + async def test_model_new_accepts_a_deployment_without_ptu_config_while_disabled(self): + _, patches = self._patched_proxy("plain-model") + admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ExitStack() as stack: + for active_patch in patches: + stack.enter_context(active_patch) + result = await add_new_model( + model_params=Deployment( + model_name="ptu-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"), + model_info=ModelInfo(id="plain-model"), + ), + user_api_key_dict=admin, + ) + + assert result.model_id == "plain-model" + + @pytest.mark.asyncio + async def test_model_new_accepts_ptu_config_once_enabled(self, monkeypatch): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + (_, add_team_model_to_db), patches = self._patched_proxy("ptu-gate-model") + admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ExitStack() as stack: + for active_patch in patches: + stack.enter_context(active_patch) + result = await add_new_model(model_params=self._ptu_deployment("ptu-gate-model"), user_api_key_dict=admin) + + assert result.model_id == "ptu-gate-model" + add_team_model_to_db.assert_called_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 1e47010b57c..6abc40eb28e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -380,6 +380,66 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth): app.dependency_overrides = {} +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"]) +@pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) +async def test_new_team_rejects_a_duration_that_never_advances( + mock_db_client, mock_admin_auth, field, bad_duration +): + """A zero-length window resets to "now", so the team row is due again the + moment it is written. The reset job re-reads such rows on every tick, and a + tenant with enough of them fills each batch and starves other tenants. + """ + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import new_team + + mock_db_client.db = MagicMock() + mock_team_create = AsyncMock() + mock_db_client.db.litellm_teamtable = MagicMock() + mock_db_client.db.litellm_teamtable.create = mock_team_create + + with pytest.raises(ProxyException) as exc_info: + await new_team( + data=NewTeamRequest(team_alias="my-team", **{field: bad_duration}), + http_request=MagicMock(spec=Request), + user_api_key_dict=mock_admin_auth, + ) + + assert str(exc_info.value.code) == "400" + assert "Invalid budget_duration" in str(exc_info.value.message) + mock_team_create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"]) +async def test_update_team_rejects_a_duration_that_never_advances( + mock_db_client, mock_admin_auth, field +): + """/team/update must reject the same never-advancing durations /team/new does.""" + from fastapi import Request + + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + mock_db_client.db = MagicMock() + mock_find_unique = AsyncMock(return_value=None) + mock_db_client.db.litellm_teamtable = MagicMock() + mock_db_client.db.litellm_teamtable.find_unique = mock_find_unique + + with pytest.raises(ProxyException) as exc_info: + await update_team( + data=UpdateTeamRequest(team_id="team-1", **{field: "0s"}), + http_request=MagicMock(spec=Request), + user_api_key_dict=mock_admin_auth, + ) + + assert str(exc_info.value.code) == "400" + assert "Invalid budget_duration" in str(exc_info.value.message) + mock_find_unique.assert_not_awaited() + + @pytest.mark.asyncio async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): """ @@ -11008,3 +11068,207 @@ def test_validate_member_user_id_provisioning_caps_the_ids_it_echoes_back(): assert f"u{_MAX_REPORTED_UNKNOWN_USER_IDS}" not in detail assert f"and {500 - _MAX_REPORTED_UNKNOWN_USER_IDS} more" in detail assert len(detail) < 1000 + + +_TEAM_ESTIMATE = "default_estimated_output_tokens" +_TEAM_ESTIMATE_PER_MODEL = "default_estimated_output_tokens_per_model" + + +@pytest.mark.parametrize( + "label, request_body, existing_metadata, allowed", + [ + ("nothing declared", {}, None, True), + ("declared top-level with none stored", {_TEAM_ESTIMATE: 1}, None, False), + ("declared inside metadata with none stored", {"metadata": {_TEAM_ESTIMATE: 1}}, None, False), + ( + "per-model map declared inside metadata", + {"metadata": {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 1}}}, + None, + False, + ), + ("unrelated edit, metadata omitted", {"tpm_limit": 99}, {_TEAM_ESTIMATE: 2000}, True), + ("stored value resent unchanged", {_TEAM_ESTIMATE: 2000}, {_TEAM_ESTIMATE: 2000}, True), + ("stored value lowered", {_TEAM_ESTIMATE: 1}, {_TEAM_ESTIMATE: 2000}, False), + ("stored value raised", {_TEAM_ESTIMATE: 9000}, {_TEAM_ESTIMATE: 2000}, False), + ( + "stored value cleared by sending a metadata blob without it", + {"metadata": {"other": "keep"}}, + {_TEAM_ESTIMATE: 2000, "other": "keep"}, + False, + ), + ( + "stored value resent inside the metadata blob", + {"metadata": {_TEAM_ESTIMATE: 2000, "other": "keep"}}, + {_TEAM_ESTIMATE: 2000, "other": "keep"}, + True, + ), + ( + "per-model map resent unchanged", + {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + True, + ), + ( + "one model in the per-model map lowered", + {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 1}}, + {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}}, + False, + ), + ], +) +def test_team_output_token_estimate_admin_gate_matrix(label, request_body, existing_metadata, allowed): + """A team admin may only leave a team's stored output-token estimate exactly as it is. + + A team admin can write team metadata, and every key on the team inherits the + team declaration, so without this a team admin could shrink the reservation + for the whole team and under-reserve against an organization TPM window the + organization set above them. Same value-transition rule as the key gate, + including the raw-metadata route and clearing by omission. + """ + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.auth.auth_utils import ( + enforce_output_token_estimates_are_admin_only, + ) + + def _call(caller): + enforce_output_token_estimates_are_admin_only( + data=UpdateTeamRequest(team_id="t", **request_body), + existing_metadata=existing_metadata, + user_api_key_dict=caller, + entity="team", + ) + + team_admin = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-admin", + user_id="team-admin", + ) + if allowed: + _call(team_admin) + else: + with pytest.raises(HTTPException) as exc: + _call(team_admin) + assert exc.value.status_code == 403 + assert "on a team" in str(exc.value.detail) + + _call( + UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + ) + + +def _wire_update_team(stack, existing_metadata): + """Mock just enough of update_team to reach (or pass) the estimate gate.""" + from unittest.mock import AsyncMock, MagicMock, patch + + mock_prisma_client = stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client")) + stack.enter_context(patch("litellm.proxy.proxy_server.llm_router")) + stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache")) + stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) + stack.enter_context(patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")) + stack.enter_context(patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object")) + + existing_team = MagicMock() + existing_team.metadata = existing_metadata + existing_team.model_dump.return_value = { + "team_id": "test_team_id", + "team_alias": "test_team", + "metadata": existing_metadata, + "members_with_roles": [{"user_id": "team-admin", "role": "admin"}], + } + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + + updated_team = MagicMock() + updated_team.team_id = "test_team_id" + updated_team.model_dump.return_value = {"team_id": "test_team_id"} + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) + mock_prisma_client.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) + return mock_prisma_client + + +@pytest.mark.asyncio +async def test_update_team_output_token_estimate_lowered_rejected_for_team_admin(): + """End-to-end wiring: _verify_team_access admits a team admin, so the gate + has to fire inside update_team itself.""" + import contextlib + from unittest.mock import Mock + + from fastapi import Request + + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + with contextlib.ExitStack() as stack: + _wire_update_team(stack, {_TEAM_ESTIMATE: 4000}) + with pytest.raises(ProxyException) as exc: + await update_team( + data=UpdateTeamRequest(team_id="test_team_id", default_estimated_output_tokens=1), + http_request=Mock(spec=Request), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-admin", + user_id="team-admin", + ), + ) + + assert str(exc.value.code) == "403" + assert "on a team" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit(): + """The team settings form resends every field it renders, so gating on + presence would break a team admin editing an unrelated setting.""" + import contextlib + from unittest.mock import Mock + + from fastapi import Request + + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + with contextlib.ExitStack() as stack: + prisma = _wire_update_team(stack, {_TEAM_ESTIMATE: 4000}) + await update_team( + data=UpdateTeamRequest( + team_id="test_team_id", + team_alias="renamed", + default_estimated_output_tokens=4000, + ), + http_request=Mock(spec=Request), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-admin", + user_id="team-admin", + ), + ) + + assert prisma.db.litellm_teamtable.update.called + + +@pytest.mark.asyncio +async def test_new_team_output_token_estimate_rejected_for_non_admin(): + """/team/new is the other write path into the same stored declaration.""" + from unittest.mock import Mock + + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import new_team + + with pytest.raises(ProxyException) as exc: + await new_team( + data=NewTeamRequest(team_alias="t", default_estimated_output_tokens=1), + http_request=Mock(spec=Request), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + ) + + assert str(exc.value.code) == "403" + assert "on a team" in str(exc.value.message) diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 91a7e1bc2c2..f70be17eb95 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -2639,3 +2639,67 @@ async def test_ProxyConfig__init_non_llm_configs_empty_agents_key_clears_remembe assert clean_agent_registry.config_agents == () clean_agent_registry.load_agents_from_db_and_config(db_agents=None) assert clean_agent_registry.get_agent_list() == () + + +# --------------------------------------------------------------------------- +# _init_guardrails_in_db +# --------------------------------------------------------------------------- + + +def _db_guardrail_row(guardrail_id: str, guardrail_type: str) -> dict[str, object]: + return { + "guardrail_id": guardrail_id, + "guardrail_name": f"name-{guardrail_id}", + "litellm_params": {"guardrail": guardrail_type, "mode": "pre_call"}, + "guardrail_info": None, + "team_id": None, + } + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_guardrails_in_db_skips_only_the_unloadable_row(monkeypatch): + """ + A single DB row that fails to initialize used to abort the whole loop, so one + typo'd guardrail type left the proxy running with zero guardrails loaded. + + The failing row's id must still reach reconcile_db_guardrails so that eviction + pass cannot treat a row that is alive in the DB as one that was deleted. + """ + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy.guardrails import guardrail_registry as registry_module + from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams + + class _RecordingHandler(registry_module.InMemoryGuardrailHandler): + def __init__(self) -> None: + super().__init__() + self.reconciled_with: list[set[str]] = [] + + def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]: + self.reconciled_with.append(set(db_guardrail_ids)) + return super().reconcile_db_guardrails(db_guardrail_ids) + + handler = _RecordingHandler() + monkeypatch.setattr(registry_module, "IN_MEMORY_GUARDRAIL_HANDLER", handler) + + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + return CustomGuardrail( + guardrail_name=guardrail["guardrail_name"], + event_hook=GuardrailEventHooks.pre_call, + default_on=False, + ) + + monkeypatch.setitem(registry_module.guardrail_initializer_registry, "lit5367_ok", _initializer) + + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[ + _db_guardrail_row("first", "lit5367_ok"), + _db_guardrail_row("broken", "litellm_tool_permission"), + _db_guardrail_row("last", "lit5367_ok"), + ] + ) + + await ProxyConfig()._init_guardrails_in_db(prisma_client=prisma_client) + + assert sorted(handler.IN_MEMORY_GUARDRAILS) == ["first", "last"] + assert handler.reconciled_with == [{"first", "broken", "last"}] diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py new file mode 100644 index 00000000000..7f4bd935a2b --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py @@ -0,0 +1,33 @@ +"""Tests for the opt-in flag that gates PTU flat-cost attribution.""" + +import pytest + +from litellm.proxy.spend_tracking.ptu_feature_flag import ( + PTU_COST_ATTRIBUTION_ENV_VAR, + is_ptu_cost_attribution_enabled, +) + + +def test_disabled_when_env_var_is_unset(monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + assert is_ptu_cost_attribution_enabled() is False + + +@pytest.mark.parametrize("value", ["true", "True", "TRUE", " true "]) +def test_enabled_for_the_values_the_house_helper_recognises(monkeypatch, value): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, value) + assert is_ptu_cost_attribution_enabled() is True + + +@pytest.mark.parametrize("value", ["false", "False", "0", "1", "", "yes", "off", "maybe"]) +def test_disabled_for_everything_else(monkeypatch, value): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, value) + assert is_ptu_cost_attribution_enabled() is False + + +def test_reads_the_env_var_on_every_call(monkeypatch): + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + assert is_ptu_cost_attribution_enabled() is False + + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + assert is_ptu_cost_attribution_enabled() is True diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index 261e3f4ef76..d17f6293cc3 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -8,6 +8,7 @@ import pytest import litellm.proxy.spend_tracking.ptu_flat_cost_rollup as ptu_rollup from litellm.constants import PTU_ROLLUP_MAX_BACKFILL_DAYS, PTU_SENTINEL_API_KEY +from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR from litellm.types.router import ModelInfo from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( PTUModel, @@ -29,6 +30,13 @@ TODAY = date(2026, 7, 31) _DEFAULT_PTU_START = "2020-01-01T00:00:00Z" +@pytest.fixture(autouse=True) +def _ptu_enabled(monkeypatch): + """PTU is gated off by default. These cover the rollup's mechanics, not the gate, so + they run with it on; the gate itself is covered by its own test below.""" + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + + _VALID_PTU = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"} @@ -1479,3 +1487,18 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts(): "a charge written 30s ago by a lagging pod was swept" ) assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-stale") not in table.rows + + +@pytest.mark.asyncio +async def test_scheduled_rollup_writes_nothing_when_ptu_attribution_is_disabled(monkeypatch): + """Startup already skips scheduling the cron, so this guards the function itself: a + deployment that never opted in accrues nothing whatever route reaches the rollup.""" + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + table = _FakeSentinelTable() + prisma = _prisma_for([_model_row(model_info=_VALID_PTU)], table) + + result = await run_scheduled_ptu_rollup(prisma, pod_lock_manager=None, alert=None) + + assert result is None + assert table.rows == {} + assert table.upsert_keys == [] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 74120beb7b9..87c8c180d9e 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11280,14 +11280,8 @@ async def test_setup_prisma_client_returns_none_when_connect_itself_fails(monkey assert mock_client.health_check.await_count == 0 -@pytest.mark.asyncio -async def test_ptu_rollup_job_registered_at_startup(monkeypatch): - """The PTU rollup cron is registered at startup; only models with PTU config accrue flat cost (asserted in test_ptu_flat_cost_rollup.py).""" - monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) +async def _run_scheduled_background_jobs(): from litellm.proxy.proxy_server import ProxyStartupEvent - from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( - PTU_ROLLUP_JOB_ID, - ) from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() @@ -11311,7 +11305,41 @@ async def test_ptu_rollup_job_registered_at_startup(monkeypatch): proxy_logging_obj=mock_proxy_logging, ) - import litellm.proxy.proxy_server as ps + import litellm.proxy.proxy_server as ps - assert ps.scheduler is not None - assert ps.scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None + assert ps.scheduler is not None + return ps.scheduler + + +@pytest.mark.asyncio +async def test_ptu_rollup_job_registered_at_startup(monkeypatch): + """The PTU rollup cron is registered once an operator opts in; only models with PTU config accrue flat cost (asserted in test_ptu_flat_cost_rollup.py).""" + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( + PTU_ROLLUP_JOB_ID, + ) + + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + + scheduler = await _run_scheduled_background_jobs() + + assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None + + +@pytest.mark.asyncio +async def test_ptu_rollup_job_not_registered_without_opt_in(monkeypatch): + """Without LITELLM_ENABLE_PTU_COST_ATTRIBUTION the rollup never runs, so no sentinel row + is ever written. This is the gate that keeps the whole feature inert by default.""" + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( + PTU_ROLLUP_JOB_ID, + ) + + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + + scheduler = await _run_scheduled_background_jobs() + + assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is None + assert len(scheduler.get_jobs()) > 0 diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 1075bffbeb2..8ee4e92ca9b 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -2928,3 +2928,143 @@ def test_update_mcp_semantic_filter_settings_requires_proxy_admin(monkeypatch): assert "proxy admin" in resp.json()["detail"].lower() finally: app.dependency_overrides.pop(user_api_key_auth, None) + + +class TestPtuCostAttributionUISetting: + """``enable_ptu_cost_attribution`` is derived from the environment on every GET. + + It is deliberately not an allowlisted, persisted setting: the point of gating PTU + flat cost on an env var is that an admin cannot flip it at runtime from the UI. + """ + + @staticmethod + def _mock_prisma(monkeypatch, stored=None): + from unittest.mock import AsyncMock, MagicMock + + mock_prisma = MagicMock() + mock_record = None + if stored is not None: + mock_record = MagicMock() + mock_record.ui_settings = stored + mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_record) + mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + return mock_prisma + + def test_reported_false_when_the_env_var_is_unset(self, mock_auth, monkeypatch): + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + self._mock_prisma(monkeypatch) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + assert response.json()["values"]["enable_ptu_cost_attribution"] is False + + def test_reported_true_once_the_env_var_is_set(self, mock_auth, monkeypatch): + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + self._mock_prisma(monkeypatch) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + assert response.json()["values"]["enable_ptu_cost_attribution"] is True + + def test_a_persisted_true_cannot_forge_the_derived_value(self, mock_auth, monkeypatch): + """A row written before the allowlist existed must not be able to turn the feature on.""" + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + self._mock_prisma(monkeypatch, stored={"enable_ptu_cost_attribution": True}) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + assert response.json()["values"]["enable_ptu_cost_attribution"] is False + + def test_is_not_an_allowlisted_persisted_setting(self): + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( + ALLOWED_UI_SETTINGS_FIELDS, + ) + + assert "enable_ptu_cost_attribution" not in ALLOWED_UI_SETTINGS_FIELDS + + def test_the_body_get_returns_is_a_valid_patch_body(self, mock_auth, monkeypatch): + """Read-modify-write is how a client edits one setting. GET injects the derived key, + so rejecting it on presence made GET's own output an invalid PATCH body: the caller + got a 400 and silently lost the edit it actually wanted.""" + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + mock_prisma = self._mock_prisma(monkeypatch) + + try: + round_tripped = client.get("/get/ui_settings").json()["values"] + assert "enable_ptu_cost_attribution" in round_tripped + response = client.patch("/update/ui_settings", json=round_tripped) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 200 + assert mock_prisma.db.litellm_uisettings.upsert.called + + def test_a_co_submitted_setting_still_applies_alongside_the_derived_key(self, mock_auth, monkeypatch): + """The derived key riding along must not discard the caller's real edit.""" + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + mock_prisma = self._mock_prisma(monkeypatch) + + try: + response = client.patch( + "/update/ui_settings", + json={"enable_ptu_cost_attribution": False, "enable_chat_ui": True}, + ) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 200 + upsert_data = mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"] + persisted = json.loads(upsert_data["create"]["ui_settings"]) + assert persisted["enable_chat_ui"] is True + assert "enable_ptu_cost_attribution" not in persisted + + def test_patch_rejects_the_derived_setting(self, mock_auth, monkeypatch): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + mock_prisma = self._mock_prisma(monkeypatch) + + try: + response = client.patch( + "/update/ui_settings", + json={"enable_ptu_cost_attribution": True}, + ) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 400 + assert "enable_ptu_cost_attribution" in str(response.json()["detail"]) + assert not mock_prisma.db.litellm_uisettings.upsert.called diff --git a/tests/test_litellm/repositories/test_unit_of_work.py b/tests/test_litellm/repositories/test_unit_of_work.py index 35f102bbb9d..c270a570ad9 100644 --- a/tests/test_litellm/repositories/test_unit_of_work.py +++ b/tests/test_litellm/repositories/test_unit_of_work.py @@ -3,7 +3,10 @@ from typing import Any, Dict, List, Mapping, Tuple import pytest -from litellm.repositories.unit_of_work import spend_reset_unit_of_work +from litellm.repositories.unit_of_work import ( + budget_cascade_unit_of_work, + spend_reset_unit_of_work, +) class FakeBatchTable: @@ -14,6 +17,9 @@ class FakeBatchTable: def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None: self._calls.append((self._table_name, dict(where), dict(data))) + def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> None: + self._calls.append((f"{self._table_name}.update_many", dict(where), dict(data))) + class FakeBatch: def __init__(self): @@ -22,6 +28,11 @@ class FakeBatch: self.litellm_verificationtoken = FakeBatchTable("litellm_verificationtoken", self.calls) self.litellm_usertable = FakeBatchTable("litellm_usertable", self.calls) self.litellm_teamtable = FakeBatchTable("litellm_teamtable", self.calls) + self.litellm_budgettable = FakeBatchTable("litellm_budgettable", self.calls) + self.litellm_teammembership = FakeBatchTable("litellm_teammembership", self.calls) + self.litellm_organizationtable = FakeBatchTable("litellm_organizationtable", self.calls) + self.litellm_tagtable = FakeBatchTable("litellm_tagtable", self.calls) + self.litellm_endusertable = FakeBatchTable("litellm_endusertable", self.calls) async def commit(self) -> None: self.commit_count += 1 @@ -64,3 +75,53 @@ async def test_empty_block_still_commits_the_batch(): assert batch.commit_count == 1 assert batch.calls == [] + + +async def test_budget_cascade_dependents_and_window_advance_share_one_batch(): + batch = FakeBatch() + reset_at = datetime(2026, 8, 3, 12, 0, tzinfo=timezone.utc) + linked = {"budget_id": {"in": ["budget-1"]}} + + async with budget_cascade_unit_of_work(lambda: batch) as uow: + uow.team_memberships.queue_spend_zero(where=linked) + uow.keys.queue_spend_zero(where=linked) + uow.organizations.queue_spend_zero(where=linked) + uow.tags.queue_spend_zero(where=linked) + uow.endusers.queue_spend_zero(where={"user_id": {"in": ["enduser-1"]}}) + uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=reset_at) + assert batch.commit_count == 0 + + assert batch.commit_count == 1 + assert batch.calls == [ + ("litellm_teammembership.update_many", linked, {"spend": 0}), + ("litellm_verificationtoken.update_many", linked, {"spend": 0}), + ("litellm_organizationtable.update_many", linked, {"spend": 0}), + ("litellm_tagtable.update_many", linked, {"spend": 0}), + ("litellm_endusertable.update_many", {"user_id": {"in": ["enduser-1"]}}, {"spend": 0}), + ("litellm_budgettable.update_many", {"budget_id": "budget-1"}, {"budget_reset_at": reset_at}), + ] + + +async def test_budget_window_advance_tolerates_a_tier_deleted_mid_chunk(): + """A tier deleted between the read and the commit must not abort the batch: + ``update`` raises P2025 on a missing row and takes every other write in the + chunk down with it, while ``update_many`` just matches nothing.""" + batch = FakeBatch() + + async with budget_cascade_unit_of_work(lambda: batch) as uow: + uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=datetime.now(timezone.utc)) + + assert [call[0] for call in batch.calls] == ["litellm_budgettable.update_many"] + + +async def test_budget_cascade_raising_inside_block_skips_commit(): + """A failure part-way through must leave budget_reset_at where it was, so + the tier is still due on the next tick.""" + batch = FakeBatch() + + with pytest.raises(RuntimeError, match="boom"): + async with budget_cascade_unit_of_work(lambda: batch) as uow: + uow.team_memberships.queue_spend_zero(where={"budget_id": {"in": ["budget-1"]}}) + raise RuntimeError("boom") + + assert batch.commit_count == 0 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 63abb2384c4..0bccb60346d 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23064 + "limit": 23057 }, "LIT002": { - "limit": 27166 + "limit": 27156 }, "LIT003": { "limit": 269 @@ -27,10 +27,10 @@ "limit": 0 }, "LIT010": { - "limit": 16753 + "limit": 16744 }, "LIT011": { - "limit": 5598 + "limit": 5596 }, "LIT012": { "limit": 5 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.test.ts new file mode 100644 index 00000000000..2215817b618 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.test.ts @@ -0,0 +1,135 @@ +import { getUiSettings } from "@/components/networking"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { renderHook, waitFor } from "@testing-library/react"; +import React, { ReactNode } from "react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { PTU_FLAG_REFRESH_MS, usePtuCostAttributionEnabled } from "./usePtuCostAttributionEnabled"; +import { useUISettings } from "./useUISettings"; + +vi.mock("@/components/networking", () => ({ + getUiSettings: vi.fn(), +})); + +describe("usePtuCostAttributionEnabled", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + vi.clearAllMocks(); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + /** Read the flag alongside the query it derives from, so assertions wait for a settled fetch. */ + const renderSettledFlag = async (settings: unknown) => { + (getUiSettings as any).mockResolvedValue(settings); + const { result } = renderHook(() => ({ enabled: usePtuCostAttributionEnabled(), query: useUISettings() }), { + wrapper, + }); + await waitFor(() => { + expect(result.current.query.isSuccess).toBe(true); + }); + return result; + }; + + it("is true only when the proxy reports the flag as enabled", async () => { + const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: true } }); + expect(result.current.enabled).toBe(true); + }); + + it("is false when the proxy reports the flag as disabled", async () => { + const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: false } }); + expect(result.current.enabled).toBe(false); + }); + + it("is false when the proxy omits the flag entirely", async () => { + const result = await renderSettledFlag({ values: { enable_chat_ui: true } }); + expect(result.current.enabled).toBe(false); + }); + + it("is false when the proxy returns no values at all", async () => { + const result = await renderSettledFlag({}); + expect(result.current.enabled).toBe(false); + }); + + it("does not treat a truthy non-boolean as enabled", async () => { + const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: "false" } }); + expect(result.current.enabled).toBe(false); + }); + + it("does not treat the string 'true' as enabled, since the proxy sends a real boolean", async () => { + const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: "true" } }); + expect(result.current.enabled).toBe(false); + }); + + it("is false before the settings request resolves", () => { + (getUiSettings as any).mockReturnValue(new Promise(() => {})); + const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper }); + expect(result.current).toBe(false); + }); + + it("is false when the settings request fails", async () => { + (getUiSettings as any).mockRejectedValue(new Error("boom")); + const { result } = renderHook(() => ({ enabled: usePtuCostAttributionEnabled(), query: useUISettings() }), { + wrapper, + }); + await waitFor(() => { + expect(result.current.query.isError).toBe(true); + }); + expect(result.current.enabled).toBe(false); + }); +}); + +describe("staleness", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + vi.clearAllMocks(); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("polls the flag once it is on, so an already-open dashboard notices it going off", async () => { + (getUiSettings as any).mockResolvedValue({ values: { enable_ptu_cost_attribution: true } }); + const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper }); + await waitFor(() => { + expect(result.current).toBe(true); + }); + + const observers = queryClient.getQueryCache().getAll()[0].observers; + const polling = observers.filter((o: any) => o.options.refetchInterval === PTU_FLAG_REFRESH_MS); + expect(polling.length).toBeGreaterThan(0); + expect(polling[0].options.staleTime).toBe(PTU_FLAG_REFRESH_MS); + expect(PTU_FLAG_REFRESH_MS).toBeLessThan(60 * 60 * 1000); + }); + + it("does not poll while the flag is off, which is every deployment that never opted in", async () => { + // The hook cannot gate on the flag before reading it, so it starts on the shared + // one-hour cache and only escalates once it has seen the feature enabled. Polling + // unconditionally made a disabled deployment re-fetch settings 120x more often. + (getUiSettings as any).mockResolvedValue({ values: { enable_ptu_cost_attribution: false } }); + const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper }); + await waitFor(() => { + expect(result.current).toBe(false); + }); + + const observers = queryClient.getQueryCache().getAll()[0].observers; + expect(observers.every((o: any) => o.options.refetchInterval === undefined)).toBe(true); + expect(observers.every((o: any) => o.options.staleTime === 60 * 60 * 1000)).toBe(true); + }); + + it("leaves the default alone for every other settings consumer", async () => { + (getUiSettings as any).mockResolvedValue({ values: {} }); + const { result } = renderHook(() => useUISettings(), { wrapper }); + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + const observers = queryClient.getQueryCache().getAll()[0].observers; + expect(observers[0].options.staleTime).toBe(60 * 60 * 1000); + expect(observers[0].options.refetchInterval).toBeUndefined(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.ts new file mode 100644 index 00000000000..e9b5afac562 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled.ts @@ -0,0 +1,26 @@ +import { useUISettings } from "./useUISettings"; + +export const PTU_COST_ATTRIBUTION_SETTING_KEY = "enable_ptu_cost_attribution"; + +/** + * Whether the proxy opted into PTU flat-cost attribution. + * + * Derived on the proxy from LITELLM_ENABLE_PTU_COST_ATTRIBUTION and returned read-only on + * /get/ui_settings, so it is not editable from the UI. Anything other than an explicit + * true (including a settings fetch that has not resolved) counts as off. + * + * Polled only once the flag has been seen on. This tracks the proxy process rather than a + * persisted setting, so an already-open dashboard has to notice a restart that turns the + * feature off, and a form that stays mounted and focused never refetches on staleTime + * alone. A deployment that never opts in is the common case and gets the shared one-hour + * cache, so the poll costs nothing where the feature is unused; the trade is that turning + * it on reaches an open dashboard on the next natural refetch rather than within 30s. + */ +export const PTU_FLAG_REFRESH_MS = 30 * 1000; + +export const usePtuCostAttributionEnabled = (): boolean => { + const { data } = useUISettings(); + const enabled = data?.values?.[PTU_COST_ATTRIBUTION_SETTING_KEY] === true; + useUISettings(enabled ? { staleTime: PTU_FLAG_REFRESH_MS, refetchInterval: PTU_FLAG_REFRESH_MS } : undefined); + return enabled; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/useUISettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/useUISettings.ts index 14c6c5e3888..749fc98c0d8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/useUISettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/uiSettings/useUISettings.ts @@ -4,11 +4,21 @@ import { createQueryKeys } from "../common/queryKeysFactory"; const uiSettingsKeys = createQueryKeys("uiSettings"); -export const useUISettings = () => { +/** + * UI settings, cached for an hour by default because they rarely change. + * + * Both options are per observer in react-query, so a caller reading a value that tracks + * proxy process state, rather than a persisted setting, can refresh it on its own cadence + * without changing how long every other caller caches. `staleTime` alone only marks the + * cached copy stale; a screen that stays mounted and focused never refetches on its own, + * so a caller that needs to notice a change also has to poll. + */ +export const useUISettings = (options?: { staleTime?: number; refetchInterval?: number }) => { return useQuery>({ queryKey: uiSettingsKeys.list({}), queryFn: async () => await getUiSettings(), - staleTime: 60 * 60 * 1000, // 1 hour - data rarely changes + staleTime: options?.staleTime ?? 60 * 60 * 1000, // 1 hour - data rarely changes gcTime: 60 * 60 * 1000, // 1 hour - keep in cache for 1 hour + refetchInterval: options?.refetchInterval, }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx index 57ff7906eda..0241ef8a77e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx @@ -25,6 +25,7 @@ import React, { useEffect, useRef, useState } from "react"; import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; import { v4 as uuidv4 } from "uuid"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import GuardrailSelector from "@/components/guardrails/GuardrailSelector"; import PolicySelector from "@/components/policies/PolicySelector"; import MCPToolArgumentsForm, { MCPToolArgumentsFormRef } from "@/components/mcp_tools/MCPToolArgumentsForm"; @@ -106,6 +107,7 @@ const ChatUI: React.FC = ({ simplified = false, fixedModel, }) => { + const canViewPolicies = useCan("viewPolicies"); const [mcpServers, setMCPServers] = useState([]); const [mcpToolsets, setMCPToolsets] = useState([]); const [isToolsetsInfoModalVisible, setIsToolsetsInfoModalVisible] = useState(false); @@ -1652,32 +1654,34 @@ const ChatUI: React.FC = ({ /> -
- - Policies - - Select policy/policies to apply to this LLM API call. Policies define which guardrails are - applied based on conditions. You can set up your policies{" "} - - here - - . - - } - > - - - - -
+ {canViewPolicies && ( +
+ + Policies + + Select policy/policies to apply to this LLM API call. Policies define which guardrails are + applied based on conditions. You can set up your policies{" "} + + here + + . + + } + > + + + + +
+ )} {/* Code Interpreter Toggle - Only for Responses endpoint */} {endpointType === EndpointType.RESPONSES && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx index 39346105f2a..c3b417987e6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx @@ -6,6 +6,7 @@ import { type ComplianceFramework, type CompliancePrompt, } from "@/data/compliancePrompts"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import { getGuardrailsList, testPoliciesAndGuardrails } from "@/components/networking"; import PolicySelector, { getPolicyOptionEntries } from "@/components/policies/PolicySelector"; import { Policy } from "@/components/policies/types"; @@ -123,6 +124,7 @@ export default function ComplianceUI({ fixedModel, proxySettings, }: ComplianceUIProps) { + const canViewPolicies = useCan("viewPolicies"); const frameworks = getFrameworks(); const [policyValueToLabel, setPolicyValueToLabel] = useState>(new Map()); @@ -701,29 +703,37 @@ export default function ComplianceUI({

Test Configuration

-

Select policies, guardrails, or both to test against.

+

+ {canViewPolicies + ? "Select policies, guardrails, or both to test against." + : "Select guardrails to test against."} +

-
- - {accessToken && ( - - )} -
+ {canViewPolicies && ( + <> +
+ + {accessToken && ( + + )} +
-
-
- or -
-
+
+
+ or +
+
+ + )}
- {/* Top Agents - only for team entity type */} - {entityType === "team" && ( + {showAgentBreakdown && (
Top Agents Driving Spend @@ -663,7 +519,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti `$${formatNumberWithCommas(value, 2)}`} @@ -685,7 +541,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti - {getProviderSpend().map((provider) => ( + {getProviderSpend(spendData.results).map((provider) => (
@@ -727,7 +583,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti ), }, - ...(entityType === "team" + ...(showAgentBreakdown ? [{ key: "agents", label: "Agent Activity", content: }] : []), { @@ -776,7 +632,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti } /> )} - {agentIsFetchingMore && entityType === "team" && ( + {agentIsFetchingMore && showAgentBreakdown && ( = ({ accessToken, entityType, enti } /> )} - {agentCancelled && entityType === "team" && ( + {agentCancelled && showAgentBreakdown && ( { + const modelSpend: { [key: string]: any } = {}; + results.forEach((day) => { + Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => { + if (!modelSpend[model]) { + modelSpend[model] = { + spend: 0, + requests: 0, + successful_requests: 0, + failed_requests: 0, + tokens: 0, + }; + } + try { + modelSpend[model].spend += metrics.metrics.spend; + } catch (e) { + console.error(`Error adding spend for ${model}: ${e}, got metrics: ${JSON.stringify(metrics)}`); + } + modelSpend[model].requests += metrics.metrics.api_requests; + modelSpend[model].successful_requests += metrics.metrics.successful_requests; + modelSpend[model].failed_requests += metrics.metrics.failed_requests; + modelSpend[model].tokens += metrics.metrics.total_tokens; + }); + }); + + return Object.entries(modelSpend) + .map(([model, metrics]) => ({ + key: model, + ...metrics, + })) + .sort((a, b) => b.spend - a.spend) + .slice(0, topModelsLimit); +}; + +export const getTopAgents = (results: ExtendedDailyData[], topAgentsLimit: number) => { + const agentSpend: { [key: string]: any } = {}; + results.forEach((day) => { + Object.entries(day.breakdown.entities || {}).forEach(([agentId, data]) => { + if (!agentSpend[agentId]) { + agentSpend[agentId] = { + spend: 0, + requests: 0, + successful_requests: 0, + failed_requests: 0, + tokens: 0, + agent_name: (data.metadata as any)?.agent_name || agentId, + }; + } + agentSpend[agentId].spend += data.metrics.spend; + agentSpend[agentId].requests += data.metrics.api_requests; + agentSpend[agentId].successful_requests += data.metrics.successful_requests; + agentSpend[agentId].failed_requests += data.metrics.failed_requests; + agentSpend[agentId].tokens += data.metrics.total_tokens; + }); + }); + + return Object.entries(agentSpend) + .map(([agentId, metrics]) => ({ + key: metrics.agent_name, + ...metrics, + })) + .sort((a, b) => b.spend - a.spend) + .slice(0, topAgentsLimit); +}; + +export const getTopAPIKeys = (results: ExtendedDailyData[], topKeysLimit: number) => { + const keySpend: { [key: string]: KeyMetricWithMetadata } = {}; + results.forEach((day) => { + const { breakdown } = day; + const { entities } = breakdown; + const tagDictionary = Object.keys(entities).reduce((acc: { [key: string]: TagUsage[] }, entity) => { + const { api_key_breakdown } = entities[entity]; + Object.keys(api_key_breakdown).forEach((key) => { + const tagUsage = { tag: entity, usage: api_key_breakdown[key].metrics.spend }; + if (acc[key]) { + acc[key].push(tagUsage); + } else { + acc[key] = [tagUsage]; + } + }); + return acc; + }, {}); + Object.entries(day.breakdown.api_keys || {}).forEach(([key, metrics]) => { + if (!keySpend[key]) { + keySpend[key] = { + metrics: { + spend: 0, + prompt_tokens: 0, + completion_tokens: 0, + total_tokens: 0, + api_requests: 0, + successful_requests: 0, + failed_requests: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }, + metadata: { + key_alias: metrics.metadata.key_alias, + team_id: metrics.metadata.team_id || null, + tags: tagDictionary[key] || [], + }, + }; + } + keySpend[key].metrics.spend += metrics.metrics.spend; + keySpend[key].metrics.prompt_tokens += metrics.metrics.prompt_tokens; + keySpend[key].metrics.completion_tokens += metrics.metrics.completion_tokens; + keySpend[key].metrics.total_tokens += metrics.metrics.total_tokens; + keySpend[key].metrics.api_requests += metrics.metrics.api_requests; + keySpend[key].metrics.successful_requests += metrics.metrics.successful_requests; + keySpend[key].metrics.failed_requests += metrics.metrics.failed_requests; + keySpend[key].metrics.cache_read_input_tokens += metrics.metrics.cache_read_input_tokens || 0; + keySpend[key].metrics.cache_creation_input_tokens += metrics.metrics.cache_creation_input_tokens || 0; + }); + }); + + return Object.entries(keySpend) + .map(([api_key, metrics]) => ({ + api_key, + key_alias: metrics.metadata.key_alias || "-", // Using truncated key as alias + tags: metrics.metadata.tags || "-", + spend: metrics.metrics.spend, + })) + .sort((a, b) => b.spend - a.spend) + .slice(0, topKeysLimit); +}; + +export const getProviderSpend = (results: ExtendedDailyData[]) => { + const providerSpend: { [key: string]: any } = {}; + results.forEach((day) => { + Object.entries(day.breakdown.providers || {}).forEach(([provider, metrics]) => { + if (!providerSpend[provider]) { + providerSpend[provider] = { + provider, + spend: 0, + requests: 0, + successful_requests: 0, + failed_requests: 0, + tokens: 0, + }; + } + try { + providerSpend[provider].spend += metrics.metrics.spend; + providerSpend[provider].requests += metrics.metrics.api_requests; + providerSpend[provider].successful_requests += metrics.metrics.successful_requests; + providerSpend[provider].failed_requests += metrics.metrics.failed_requests; + providerSpend[provider].tokens += metrics.metrics.total_tokens; + } catch (e) { + console.error(`Error processing provider ${provider}: ${e}`); + } + }); + }); + + return Object.values(providerSpend) + .filter((provider) => provider.spend > 0) + .sort((a, b) => b.spend - a.spend); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx index 0ded7d195d0..9085cf961a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx @@ -502,6 +502,8 @@ describe("UsagePage", () => { userId: "user-123", userEmail: "test@example.com", userRole: "Internal User", + userRoleLabel: "Internal User", + isViewOnly: false, premiumUser: true, disabledPersonalKeyCreation: false, showSSOBanner: false, @@ -861,6 +863,27 @@ describe("UsagePage", () => { }); }); + it.each(["organization", "agent"])("should not render the %s usage view for an internal user", async (usageView) => { + mockUseAuthorized.mockReturnValue(nonAdminSession); + + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + const usageSelect = screen.getByTestId("usage-view-select"); + act(() => { + fireEvent.change(usageSelect, { target: { value: "team" } }); + }); + expect(screen.getAllByText("Entity Usage").length).toBeGreaterThan(0); + + act(() => { + fireEvent.change(usageSelect, { target: { value: usageView } }); + }); + expect(screen.queryByText("Entity Usage")).not.toBeInTheDocument(); + }); + describe("admin user selector", () => { it("should render user selector for admin users in global view", async () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index c3645d6371e..494df313ac0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -33,6 +33,7 @@ import { useCustomers } from "@/app/(dashboard)/hooks/customers/useCustomers"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser"; import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers"; +import { hasCapability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { all_admin_roles, internalUserRoles } from "@/utils/roles"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; @@ -109,6 +110,8 @@ const UsagePage: React.FC = ({ teams, organizations }) => { const { data: currentUser } = useCurrentUser(); const isAdmin = all_admin_roles.includes(userRole || ""); const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || ""); + const canViewOrganizationUsage = hasCapability(userRole, "viewOrganizationUsage"); + const canViewAgentUsage = hasCapability(userRole, "viewAgentUsage"); // Debounced search for user selector const [userSearchInput, setUserSearchInput] = useState(""); @@ -513,7 +516,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { setUsageView(value)} - isAdmin={isAdmin} + userRole={userRole} canViewTagUsage={canViewTagUsage} /> @@ -950,7 +953,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { )} {/* Organization Usage Panel */} - {usageView === "organization" && ( + {usageView === "organization" && canViewOrganizationUsage && ( = ({ teams, organizations }) => { /> )} - {usageView === "agent" && ( + {usageView === "agent" && canViewAgentUsage && ( { }); it("should render", () => { - render(); + render(); expect(screen.getByText("Usage View")).toBeInTheDocument(); expect(screen.getByText("Select the usage data you want to view")).toBeInTheDocument(); expect(screen.getByRole("combobox")).toBeInTheDocument(); + expect(screen.getByRole("option", { name: "Your Usage" })).toBeInTheDocument(); }); it("should call onChange when value changes", () => { - render(); + render(); const select = screen.getByRole("combobox"); act(() => { @@ -109,14 +110,32 @@ describe("UsageViewSelect", () => { }); it("should show Tag Usage for non-admin users with tag usage permission", () => { - render(); + render(); expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument(); }); it("should hide Tag Usage for non-admin users without tag usage permission", () => { - render(); + render(); expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument(); }); + + it.each(["Organization Usage", "Agent Usage (A2A)"])("should show %s to an admin", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); + + it.each(["Organization Usage", "Agent Usage (A2A)"])("should hide %s from an internal user", (optionName) => { + render(); + + expect(screen.queryByRole("option", { name: optionName })).not.toBeInTheDocument(); + }); + + it.each(["Team Usage", "Tag Usage"])("should keep %s available to an internal user", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx index 94b483cb539..54c1d5ab7cc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx @@ -11,6 +11,8 @@ import { } from "@ant-design/icons"; import { Badge, Select } from "antd"; import React from "react"; +import { hasCapability, type Capability } from "@/utils/capabilities"; +import { all_admin_roles } from "@/utils/roles"; export type UsageOption = | "global" | "my-usage" @@ -24,7 +26,7 @@ export type UsageOption = export interface UsageViewSelectProps { value: UsageOption; onChange: (value: UsageOption) => void; - isAdmin: boolean; + userRole: string | null; canViewTagUsage?: boolean; title?: string; description?: string; @@ -35,6 +37,7 @@ interface OptionConfig { label: string; description: string; icon: React.ReactNode; + capability?: Capability; adminOnly?: boolean; showForAdmin?: string; showForNonAdmin?: string; @@ -63,12 +66,9 @@ const OPTIONS: OptionConfig[] = [ { value: "organization", label: "Organization Usage", - showForAdmin: "Organization Usage", - showForNonAdmin: "Your Organization Usage", - description: "View organization-level usage", - descriptionForAdmin: "View usage across all organizations", - descriptionForNonAdmin: "View your organization's usage", + description: "View usage across all organizations", icon: , + capability: "viewOrganizationUsage", }, { value: "team", @@ -95,7 +95,7 @@ const OPTIONS: OptionConfig[] = [ label: "Agent Usage (A2A)", description: "View usage by AI agents", icon: , - adminOnly: true, + capability: "viewAgentUsage", }, { value: "user", @@ -115,14 +115,18 @@ const OPTIONS: OptionConfig[] = [ export const UsageViewSelect: React.FC = ({ value, onChange, - isAdmin, + userRole, canViewTagUsage = false, title = "Usage View", description = "Select the usage data you want to view", "data-id": dataId, }) => { + const isAdmin = all_admin_roles.includes(userRole ?? ""); const getFilteredOptions = () => { return OPTIONS.filter((option) => { + if (option.capability) { + return hasCapability(userRole, option.capability); + } if (option.value === "tag" && canViewTagUsage) { return true; } diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index e8331294972..c5c77c0c878 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -6,9 +6,14 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; import { useTeamMetadataSchema } from "@/app/(dashboard)/hooks/teams/useTeamMetadataSchema"; import NotificationsManager from "./molecules/notifications_manager"; import { fetchAvailableModelsForTeamOrKey } from "./key_team_helpers/fetch_available_models_team_key"; -import { fetchMCPAccessGroups, getGuardrailsList, teamCreateCall } from "./networking"; +import { fetchMCPAccessGroups, getGuardrailsList, getPoliciesList, teamCreateCall } from "./networking"; import Teams from "./Teams"; +const can = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useCan", () => ({ + default: (...args: unknown[]) => can(...args), +})); + const mockTeamInfoView = vi.fn(); const mockUseOrganizations = vi.fn(); @@ -173,6 +178,7 @@ const renderWithQueryClient = ( // Re-establish safe defaults before every test (clearAllMocks keeps return values, so restore them here). beforeEach(() => { mockTeamsTableProps = null; + can.mockReturnValue(true); }); describe("Teams - handleCreate organization handling", () => { @@ -956,3 +962,50 @@ describe("Teams - LIT-2530 organization stays optional for proxy admin with a si }); }); }); + +describe("Teams - policies field is gated on the viewPolicies capability", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockTeamInfoView.mockClear(); + vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4"]); + vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]); + vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] }); + vi.mocked(getPoliciesList).mockResolvedValue({ policies: [] }); + mockUseOrganizations.mockReturnValue({ data: null }); + }); + + const openAdditionalSettings = async () => { + renderWithQueryClient(); + + act(() => { + fireEvent.click(screen.getAllByRole("button", { name: /create team/i })[0]); + }); + + await waitFor(() => { + expect(screen.getByLabelText(/team name/i)).toBeInTheDocument(); + }); + + fireEvent.click(screen.getByText("Additional Settings")); + + await waitFor(() => { + expect(screen.getByTestId("access-group-selector")).toBeInTheDocument(); + }); + }; + + it("should render the policies field and load it when the capability is present", async () => { + await openAdditionalSettings(); + + expect(can).toHaveBeenCalledWith("viewPolicies"); + expect(getPoliciesList).toHaveBeenCalledWith("test-token"); + expect(screen.getByText("Policies")).toBeInTheDocument(); + }); + + it("should omit the policies field and skip the admin-only list without the capability", async () => { + can.mockReturnValue(false); + + await openAdditionalSettings(); + + expect(getPoliciesList).not.toHaveBeenCalled(); + expect(screen.queryByText("Policies")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index 42ff8bcc4c2..7c269555ec0 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -1,4 +1,5 @@ import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import AvailableTeamsPanel from "@/components/team/AvailableTeamsPanel"; import TeamInfoView from "@/components/team/TeamInfo"; import TeamSSOSettings from "@/components/TeamSSOSettings"; @@ -108,6 +109,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser const [isTeamDeleting, setIsTeamDeleting] = useState(false); // Add this state near the other useState declarations const [guardrailsList, setGuardrailsList] = useState([]); + const canViewPolicies = useCan("viewPolicies"); const [policiesList, setPoliciesList] = useState([]); const [loggingSettings, setLoggingSettings] = useState([]); const [modelAliases, setModelAliases] = useState<{ [key: string]: string }>({}); @@ -168,8 +170,8 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser }; fetchGuardrails(); - fetchPolicies(); - }, [accessToken]); + if (canViewPolicies) fetchPolicies(); + }, [accessToken, canViewPolicies]); const handleOk = () => { setIsTeamModalVisible(false); @@ -795,36 +797,38 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser } /> - - Policies{" "} - - e.stopPropagation()} - > - - - - - } - name="policies" - className="mt-8" - help="Select existing policies or enter new ones" - > - ({ + value: name, + label: name, + }))} + /> + + )} diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx index 9fe36e13998..4e5c5f25374 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx @@ -2,32 +2,37 @@ import { act, fireEvent, render, waitFor } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import AdvancedSettings from "./advanced_settings"; +const mockUsePtuCostAttributionEnabled = vi.fn(); + +vi.mock("@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled", () => ({ + usePtuCostAttributionEnabled: () => mockUsePtuCostAttributionEnabled(), +})); + +const PTU_LABELS = ["PTU Count", "Calculated Cost per PTU / Hour (USD)", "PTU Effective From (UTC)"]; + +const renderAdvancedSettings = () => + render( + {}} + guardrailsList={[]} + tagsList={{}} + accessToken="test-token" + />, + ); + describe("AdvancedSettings", () => { beforeEach(() => { vi.clearAllMocks(); + mockUsePtuCostAttributionEnabled.mockReturnValue(false); }); + it("should render", () => { - render( - {}} - guardrailsList={[]} - tagsList={{}} - accessToken="test-token" - />, - ); + renderAdvancedSettings(); }); it("should render tags list", async () => { - const { getByText } = render( - {}} - guardrailsList={[]} - tagsList={{}} - accessToken="test-token" - />, - ); + const { getByText } = renderAdvancedSettings(); fireEvent.click(getByText("Advanced Settings")); await waitFor(() => { expect(getByText("Tags")).toBeInTheDocument(); @@ -35,15 +40,7 @@ describe("AdvancedSettings", () => { }); it("should render the litellm params", async () => { - const { getByText } = render( - {}} - guardrailsList={[]} - tagsList={{}} - accessToken="test-token" - />, - ); + const { getByText } = renderAdvancedSettings(); act(() => { fireEvent.click(getByText("Advanced Settings")); }); @@ -51,4 +48,35 @@ describe("AdvancedSettings", () => { expect(getByText("LiteLLM Params")).toBeInTheDocument(); }); }); + + it("hides every PTU field when PTU cost attribution is disabled", async () => { + const { getByText, queryByText } = renderAdvancedSettings(); + act(() => { + fireEvent.click(getByText("Advanced Settings")); + }); + await waitFor(() => { + expect(getByText("Tags")).toBeInTheDocument(); + }); + + for (const label of PTU_LABELS) { + expect(queryByText(label)).not.toBeInTheDocument(); + } + expect(queryByText("PTU Effective To (UTC)")).not.toBeInTheDocument(); + }); + + it("shows every PTU field when PTU cost attribution is enabled", async () => { + mockUsePtuCostAttributionEnabled.mockReturnValue(true); + const { getByText } = renderAdvancedSettings(); + act(() => { + fireEvent.click(getByText("Advanced Settings")); + }); + + await waitFor(() => { + expect(getByText("PTU Count")).toBeInTheDocument(); + }); + for (const label of PTU_LABELS) { + expect(getByText(label)).toBeInTheDocument(); + } + expect(getByText("PTU Effective To (UTC)")).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx index 9bc78783a57..5d7196a84fe 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx @@ -20,6 +20,7 @@ import { ptuWindowOrderRule, PTU_END_FIELD, } from "../../utils/ptuValidation"; +import { usePtuCostAttributionEnabled } from "@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled"; const { Link } = Typography; interface AdvancedSettingsProps { @@ -43,6 +44,7 @@ const AdvancedSettings: React.FC = ({ const [customPricing, setCustomPricing] = React.useState(false); const [pricingModel, setPricingModel] = React.useState<"per_token" | "per_second">("per_token"); const [showCacheControl, setShowCacheControl] = React.useState(false); + const ptuCostAttributionEnabled = usePtuCostAttributionEnabled(); // Add validation function for numbers const validateNumber = (_: any, value: string) => { @@ -193,49 +195,53 @@ const AdvancedSettings: React.FC = ({ /> - - - + {ptuCostAttributionEnabled && ( + <> + + + - - - + + + - - - + + + - - - + + + + + )} {customPricing && (
diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index a5b273a0f56..f795076ff03 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -210,6 +210,7 @@ describe("Sidebar (leftnav)", () => { userId: "internal-user-id", accessToken: "test-access-token", userRole: "internal", + isViewOnly: false, token: "test-token", userEmail: "internal@example.com", premiumUser: false, @@ -244,6 +245,27 @@ describe("Sidebar (leftnav)", () => { expect(screen.getByText("Tool Policies")).toBeInTheDocument(); }); }); + + it("should hide the Policies entry from internal users while keeping Guardrails", () => { + mockUseAuthorized.mockReturnValue(internalAuth); + renderWithProviders(); + + expect(screen.getByText("Guardrails")).toBeInTheDocument(); + expect(screen.queryByText("Policies")).not.toBeInTheDocument(); + }); + + it("should hide the Prompts entry from internal users while keeping other Experimental children", async () => { + mockUseAuthorized.mockReturnValue(internalAuth); + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByText("Experimental")); + }); + await waitFor(() => { + expect(screen.getByText("API Playground")).toBeInTheDocument(); + }); + expect(screen.queryByText("Prompts")).not.toBeInTheDocument(); + }); }); it("should show Organizations tab for organization admins", () => { diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index f08092d0e38..12d124b059e 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -158,7 +158,7 @@ const menuGroups: MenuGroup[] = [ page: "policies", label: "Policies", icon: , - roles: all_admin_roles, + roles: rolesWithCapability("viewPolicies"), }, { key: "tools", @@ -268,7 +268,13 @@ const menuGroups: MenuGroup[] = [ label: "Experimental", icon: , children: [ - { key: "prompts", page: "prompts", label: "Prompts", icon: , roles: all_admin_roles }, + { + key: "prompts", + page: "prompts", + label: "Prompts", + icon: , + roles: rolesWithCapability("viewPrompts"), + }, { key: "transform-request", page: "transform-request", diff --git a/ui/litellm-dashboard/src/components/model_info_view.test.tsx b/ui/litellm-dashboard/src/components/model_info_view.test.tsx index 326b49ff896..a3ce40494eb 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.test.tsx @@ -47,6 +47,11 @@ vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ useModelCostMap: (...args: any[]) => mockUseModelCostMap(...args), })); +const mockUsePtuCostAttributionEnabled = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled", () => ({ + usePtuCostAttributionEnabled: () => mockUsePtuCostAttributionEnabled(), +})); + const mockNotificationsManager = vi.mocked(NotificationsManager); const mockModelInfoV1Call = vi.mocked(networking.modelInfoV1Call); const mockCredentialGetCall = vi.mocked(networking.credentialGetCall); @@ -99,6 +104,7 @@ describe("ModelInfoView", () => { }, }); vi.clearAllMocks(); + mockUsePtuCostAttributionEnabled.mockReturnValue(false); mockUseModelsInfo.mockReturnValue({ data: { @@ -608,6 +614,100 @@ describe("ModelInfoView", () => { expect(updatePayload.litellm_params).not.toHaveProperty("vector_store_ids"); }); + describe("PTU cost attribution gate", () => { + const ptuModelData = { + ...defaultModelData, + model_info: { + ...defaultModelData.model_info, + team_id: "team-1", + ptu_count: 15, + cost_per_ptu_per_hour: 2, + ptu_effective_from: "2026-07-01T00:00:00+00:00", + ptu_effective_to: "2026-08-01T00:00:00+00:00", + }, + }; + + const renderWithPtuModel = () => { + mockUseModelsInfo.mockReturnValue({ data: { data: [ptuModelData] }, isLoading: false, error: null }); + mockModelInfoV1Call.mockResolvedValue({ data: [ptuModelData] }); + return render(, { wrapper }); + }; + + it("hides the PTU fields when disabled, even for a model that already stores PTU config", async () => { + renderWithPtuModel(); + + await waitFor(() => { + expect(screen.getByText("Model Settings")).toBeInTheDocument(); + }); + + expect(screen.queryByText("PTU Count")).not.toBeInTheDocument(); + expect(screen.queryByText("Cost per PTU / Hour (USD)")).not.toBeInTheDocument(); + expect(screen.queryByText("PTU Effective From (UTC)")).not.toBeInTheDocument(); + expect(screen.queryByText("PTU Effective To (UTC)")).not.toBeInTheDocument(); + }); + + it("shows the PTU fields when enabled", async () => { + mockUsePtuCostAttributionEnabled.mockReturnValue(true); + renderWithPtuModel(); + + await waitFor(() => { + expect(screen.getByText("PTU Count")).toBeInTheDocument(); + }); + expect(screen.getByText("Cost per PTU / Hour (USD)")).toBeInTheDocument(); + expect(screen.getByText("PTU Effective From (UTC)")).toBeInTheDocument(); + expect(screen.getByText("PTU Effective To (UTC)")).toBeInTheDocument(); + }); + + it("omits PTU fields from the save payload when disabled, so an unrelated edit cannot clear stored config", async () => { + const user = userEvent.setup(); + renderWithPtuModel(); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); + }); + await user.click(screen.getByRole("button", { name: /edit settings/i })); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(mockModelPatchUpdateCall).toHaveBeenCalled(); + }); + + const modelInfo = mockModelPatchUpdateCall.mock.calls[0][1].model_info; + expect(modelInfo).not.toHaveProperty("ptu_count"); + expect(modelInfo).not.toHaveProperty("cost_per_ptu_per_hour"); + expect(modelInfo).not.toHaveProperty("ptu_effective_from"); + expect(modelInfo).not.toHaveProperty("ptu_effective_to"); + }); + + it("sends the PTU fields on save when enabled", async () => { + mockUsePtuCostAttributionEnabled.mockReturnValue(true); + const user = userEvent.setup(); + renderWithPtuModel(); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); + }); + await user.click(screen.getByRole("button", { name: /edit settings/i })); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(mockModelPatchUpdateCall).toHaveBeenCalled(); + }); + + const modelInfo = mockModelPatchUpdateCall.mock.calls[0][1].model_info; + expect(modelInfo.ptu_count).toBe(15); + expect(modelInfo.cost_per_ptu_per_hour).toBe(2); + }); + }); + it("should not include input_cost_per_token or output_cost_per_token in update payload when user does not touch cost fields", async () => { // Regression: editing a model without touching cost fields used to inject // input_cost_per_token: 0 and output_cost_per_token: 0 into litellm_params, diff --git a/ui/litellm-dashboard/src/components/model_info_view.tsx b/ui/litellm-dashboard/src/components/model_info_view.tsx index bd8da531a8f..e1a4d4311ab 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.tsx @@ -18,7 +18,9 @@ import { Button as TremorButton, } from "@tremor/react"; import { Button, DatePicker, Form, Input, Modal, Select, Tooltip } from "antd"; -import { formatPtuUtcDisplay, ptuPickerToUtcIso, utcIsoToPickerValue } from "../utils/ptuDatetime"; +import { formatPtuUtcDisplay, utcIsoToPickerValue } from "../utils/ptuDatetime"; +import { applyPtuModelInfo } from "../utils/ptuModelInfo"; +import { usePtuCostAttributionEnabled } from "@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled"; import { PTU_COUNT_FIELD, PTU_RATE_FIELD, @@ -224,6 +226,7 @@ export default function ModelInfoView({ const { data: modelCostMapData } = useModelCostMap(); const { data: modelHubData } = useModelHub(); const { data: teams } = useTeams(); + const ptuCostAttributionEnabled = usePtuCostAttributionEnabled(); // Transform the model data const getProviderFromModel = (model: string) => { @@ -495,15 +498,7 @@ export default function ModelInfoView({ health_check_model: values.health_check_model, }; } - const ptuNumber = (val: string | number | null | undefined): number | null => - val !== undefined && val !== null && val !== "" ? Number(val) : null; - updatedModelInfo = { - ...updatedModelInfo, - ptu_count: ptuNumber(values.ptu_count), - cost_per_ptu_per_hour: ptuNumber(values.cost_per_ptu_per_hour), - ptu_effective_from: ptuPickerToUtcIso(values.ptu_effective_from), - ptu_effective_to: ptuPickerToUtcIso(values.ptu_effective_to), - }; + updatedModelInfo = applyPtuModelInfo(updatedModelInfo, values, ptuCostAttributionEnabled); } catch (e) { NotificationsManager.fromBackend("Invalid JSON in Model Info"); return; @@ -953,45 +948,46 @@ export default function ModelInfoView({ )}
- {PTU_EDIT_FIELDS.map((ptuField) => { - const { name, label, input, placeholder, isCount, isRate, isStart, pairedWith } = ptuField; - const { windowPeer, bound } = ptuField; - return ( -
- {label} - {isEditing ? ( - - {input === "number" ? ( - - ) : ( - - )} - - ) : ( -
- {(input === "datetime" - ? formatPtuUtcDisplay(localModelData?.model_info?.[name]) - : localModelData?.model_info?.[name]) ?? "Not Set"} -
- )} -
- ); - })} + {ptuCostAttributionEnabled && + PTU_EDIT_FIELDS.map((ptuField) => { + const { name, label, input, placeholder, isCount, isRate, isStart, pairedWith } = ptuField; + const { windowPeer, bound } = ptuField; + return ( +
+ {label} + {isEditing ? ( + + {input === "number" ? ( + + ) : ( + + )} + + ) : ( +
+ {(input === "datetime" + ? formatPtuUtcDisplay(localModelData?.model_info?.[name]) + : localModelData?.model_info?.[name]) ?? "Not Set"} +
+ )} +
+ ); + })}
Cache Read Cost (per 1M tokens) diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx index 2f72fa37d86..5071fe6ff95 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx @@ -2,7 +2,7 @@ import { act, fireEvent, within } from "@testing-library/react"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; import { Team } from "../key_team_helpers/key_list"; -import { userFilterUICall } from "../networking"; +import { getPoliciesList, getPromptsList, userFilterUICall } from "../networking"; import CreateKey from "./create_key_button"; const { formMock, setFieldsValueMock, radioGroupValueRef, formStateRef, mockKeyCreateCall, teamDropdownTeamsRef } = @@ -777,4 +777,47 @@ describe("CreateKey", () => { }); }); }); + + describe("policy and prompt fields", () => { + const POLICIES_PLACEHOLDER = "Premium feature - Upgrade to set policies by key"; + const PROMPTS_PLACEHOLDER = "Premium feature - Upgrade to set prompts by key"; + + const openModal = () => { + renderWithProviders(); + act(() => { + fireEvent.click(screen.getByRole("button", { name: /create new key/i })); + }); + }; + + beforeEach(() => { + vi.mocked(getPoliciesList).mockResolvedValue({ policies: [{ policy_name: "policy-a" }] }); + vi.mocked(getPromptsList).mockResolvedValue({ prompts: [{ prompt_id: "prompt-a" }] } as any); + }); + + it("should load and offer both selectors for an admin", async () => { + openModal(); + + await waitFor(() => { + expect(screen.getByRole("option", { name: "policy-a" })).toBeInTheDocument(); + expect(screen.getByRole("option", { name: "prompt-a" })).toBeInTheDocument(); + }); + expect(getPoliciesList).toHaveBeenCalledWith("test-token"); + expect(getPromptsList).toHaveBeenCalledWith("test-token"); + expect(screen.getByPlaceholderText(POLICIES_PLACEHOLDER)).toBeInTheDocument(); + expect(screen.getByPlaceholderText(PROMPTS_PLACEHOLDER)).toBeInTheDocument(); + }); + + it("should omit both selectors and fire neither admin-only request for an internal user", async () => { + authorizedState = { ...defaultAuthorizedState, userRole: "Internal User" }; + + openModal(); + + expect(await screen.findByTestId("org-dropdown")).toBeInTheDocument(); + + expect(getPoliciesList).not.toHaveBeenCalled(); + expect(getPromptsList).not.toHaveBeenCalled(); + expect(screen.queryByPlaceholderText(POLICIES_PLACEHOLDER)).not.toBeInTheDocument(); + expect(screen.queryByPlaceholderText(PROMPTS_PLACEHOLDER)).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 189df662cd8..19abae73549 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -5,6 +5,7 @@ import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import { useTags } from "@/app/(dashboard)/hooks/tags/useTags"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { InfoCircleOutlined } from "@ant-design/icons"; import { useQueryClient } from "@tanstack/react-query"; @@ -147,6 +148,8 @@ export const fetchUserModels = async ( const CreateKey: React.FC = ({ team, teams, data, addKey, autoOpenCreate, prefillData }) => { const { accessToken, userId: userID, userRole, premiumUser } = useAuthorized(); const canEditGuardrails = premiumUser || (userRole != null && rolesWithWriteAccess.includes(userRole)); + const canViewPolicies = useCan("viewPolicies"); + const canViewPrompts = useCan("viewPrompts"); const { data: organizations, isLoading: isOrganizationsLoading } = useOrganizations(); const { data: projects, isLoading: isProjectsLoading } = useProjects(); const { data: uiSettingsData } = useUISettings(); @@ -275,9 +278,9 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp }; fetchGuardrails(); - fetchPolicies(); - fetchPrompts(); - }, [accessToken]); + if (canViewPolicies) fetchPolicies(); + if (canViewPrompts) fetchPrompts(); + }, [accessToken, canViewPolicies, canViewPrompts]); // Fetch possible user roles when component mounts useEffect(() => { @@ -1251,74 +1254,78 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp > - - Policies{" "} - - e.stopPropagation()} // Prevent accordion from collapsing when clicking link - > - - - - - } - name="policies" - className="mt-4" - help={ - premiumUser - ? "Select existing policies or enter new ones" - : "Premium feature - Upgrade to set policies by key" - } - > - ({ value: name, label: name }))} - /> - + > + ({ value: name, label: name }))} + /> + + )} diff --git a/ui/litellm-dashboard/src/components/policies/PolicySelector.test.tsx b/ui/litellm-dashboard/src/components/policies/PolicySelector.test.tsx index 7a263743688..84dfb5cc1c5 100644 --- a/ui/litellm-dashboard/src/components/policies/PolicySelector.test.tsx +++ b/ui/litellm-dashboard/src/components/policies/PolicySelector.test.tsx @@ -7,6 +7,11 @@ import { Policy } from "./types"; vi.mock("../networking"); +const can = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useCan", () => ({ + default: (...args: unknown[]) => can(...args), +})); + const makePolicy = (overrides: Partial): Policy => ({ policy_id: "uuid-1", policy_name: "test-policy", @@ -76,6 +81,7 @@ describe("PolicySelector", () => { beforeEach(() => { vi.clearAllMocks(); + can.mockReturnValue(true); }); it("should render", () => { @@ -114,4 +120,18 @@ describe("PolicySelector", () => { renderWithProviders(); expect(networking.getPoliciesList).not.toHaveBeenCalled(); }); + + it("should render nothing and skip the admin-only fetch without the viewPolicies capability", async () => { + can.mockReturnValue(false); + vi.mocked(networking.getPoliciesList).mockResolvedValue({ policies: [] }); + + const { container } = renderWithProviders(); + + await waitFor(() => { + expect(can).toHaveBeenCalledWith("viewPolicies"); + }); + expect(networking.getPoliciesList).not.toHaveBeenCalled(); + expect(screen.queryByRole("combobox")).not.toBeInTheDocument(); + expect(container).toBeEmptyDOMElement(); + }); }); diff --git a/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx b/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx index 132538d439f..69816fafc50 100644 --- a/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx +++ b/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx @@ -1,5 +1,6 @@ import React, { useEffect, useState } from "react"; import { Select } from "antd"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import { Policy } from "./types"; import { getPoliciesList } from "../networking"; @@ -51,12 +52,13 @@ const PolicySelector: React.FC = ({ disabled, onPoliciesLoaded, }) => { + const canViewPolicies = useCan("viewPolicies"); const [policies, setPolicies] = useState([]); const [loading, setLoading] = useState(false); useEffect(() => { const fetchPolicies = async () => { - if (!accessToken) return; + if (!accessToken || !canViewPolicies) return; setLoading(true); try { @@ -73,12 +75,16 @@ const PolicySelector: React.FC = ({ }; fetchPolicies(); - }, [accessToken, onPoliciesLoaded]); + }, [accessToken, canViewPolicies, onPoliciesLoaded]); const handlePolicyChange = (selectedValues: string[]) => { onChange(selectedValues); }; + if (!canViewPolicies) { + return null; + } + return (
({ value: name, label: name }))} - /> - + {canViewPolicies && ( + + Policies{" "} + + e.stopPropagation()} + > + + + + + } + name="policies" + > + + + + + + + + + @@ -610,27 +602,29 @@ export function KeyEditView({ - - Policies{" "} - - - - - } - name="policies" - > - {accessToken && ( - { - form.setFieldValue("policies", v); - }} - accessToken={accessToken} - disabled={!premiumUser} - /> - )} - + {canViewPolicies && ( + + Policies{" "} + + + + + } + name="policies" + > + {accessToken && ( + { + form.setFieldValue("policies", v); + }} + accessToken={accessToken} + disabled={!premiumUser} + /> + )} + + )} 0 - ? `Current: ${keyData.metadata.prompts.join(", ")}` - : "Select or enter prompts" - } - options={promptsList.map((name) => ({ value: name, label: name }))} - /> - - + {canViewPrompts && ( + + +