merge: litellm_internal_staging into litellm_slot_leak_stream_logging

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-08-20 05:25:18 +00:00
commit a1b32e446d
425 changed files with 9919 additions and 6952 deletions

View file

@ -1,15 +1,17 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client>}"
category="${1:?usage: classify_changes.sh <backend|client|ui>}"
has_client=false
has_backend=false
has_ci=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
.github/* | .circleci/*) has_ci=true; has_backend=true ;;
*) has_backend=true ;;
esac
done
@ -21,6 +23,9 @@ case "$category" in
client)
{ [ "$has_client" = true ] || [ "$has_backend" = true ]; } && echo run || echo skip
;;
ui)
{ [ "$has_client" = true ] || [ "$has_ci" = true ]; } && echo run || echo skip
;;
*)
echo run
;;

View file

@ -1,34 +0,0 @@
name: "Detect backend-relevant changes"
description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip. decision=skip means only ui/**, **.md or **.mdx files
changed, so callers can short-circuit expensive steps while the job still completes
successfully and satisfies its required status check. The file list comes from the
pull request itself rather than from a git diff, because the checked-out merge ref is
recomputed as the base branch advances and would otherwise attribute the base
branch's own commits to the pull request. The decision defaults to run for any non
pull_request event or whenever the changed set cannot be resolved, so tests are never
skipped when the classification is uncertain.
inputs:
github-token:
description: "Token used to list the pull request's files; needs pull-requests: read"
required: false
default: ${{ github.token }}
outputs:
decision:
description: "run when backend-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
GH_TOKEN: ${{ inputs.github-token }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
CHANGED_FILE_COUNT: ${{ github.event.pull_request.changed_files }}
run: bash "${GITHUB_ACTION_PATH}/../../scripts/detect_backend_changes.sh"

View file

@ -0,0 +1,41 @@
name: "Detect relevant changes"
description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip for one category. backend means anything outside ui/,
docs/ and markdown; ui means the dashboard sources alone. decision=skip lets callers
short-circuit expensive steps while the job still completes successfully and satisfies
its required status check, which a paths: filter cannot do because a workflow that
never starts never reports. The file list comes from the pull request itself rather
than from a git diff, because the checked-out merge ref is recomputed as the base
branch advances and would otherwise attribute the base branch's own commits to the
pull request. The decision defaults to run for any non pull_request event or whenever
the changed set cannot be resolved, so jobs are never skipped when the classification
is uncertain.
inputs:
category:
description: "Which classification to apply: backend, client or ui"
required: false
default: backend
github-token:
description: "Token used to list the pull request's files; needs pull-requests: read"
required: false
default: ${{ github.token }}
outputs:
decision:
description: "run when category-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
GH_TOKEN: ${{ inputs.github-token }}
CATEGORY: ${{ inputs.category }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
CHANGED_FILE_COUNT: ${{ github.event.pull_request.changed_files }}
run: bash "${GITHUB_ACTION_PATH}/../../scripts/detect_changes.sh"

View file

@ -2,15 +2,16 @@
set -uo pipefail
readonly API_FILE_CEILING=3000
readonly CATEGORY="${CATEGORY:-backend}"
decide() {
echo "detect-backend-changes: decision=$1"
echo "detect-changes[${CATEGORY}]: decision=$1"
[ -z "${GITHUB_OUTPUT:-}" ] || echo "decision=$1" >>"${GITHUB_OUTPUT}"
exit 0
}
run_full() {
echo "detect-backend-changes: $1; running job"
echo "detect-changes[${CATEGORY}]: $1; running job"
decide run
}
@ -30,10 +31,10 @@ changed="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate --jq '.[].
run_full "could not list the files on PR #${PR_NUMBER}"
[ -n "${changed}" ] || run_full "the API listed no files on PR #${PR_NUMBER}"
echo "detect-backend-changes: files changed by PR #${PR_NUMBER}:"
echo "detect-changes[${CATEGORY}]: files changed by PR #${PR_NUMBER}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash "${classify}" backend)" ||
decision="$(printf '%s\n' "${changed}" | bash "${classify}" "${CATEGORY}")" ||
run_full "classify_changes.sh failed"
case "${decision}" in
run | skip) decide "${decision}" ;;

View file

@ -72,24 +72,27 @@ jobs:
with:
persist-credentials: false
- name: Detect backend-relevant changes
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-backend-changes
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:

View file

@ -24,6 +24,7 @@ jobs:
# re-running basedpyright over the merge-base tree.
permissions:
contents: read
pull-requests: read
actions: read
steps:
@ -37,7 +38,12 @@ jobs:
clean: true
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Fetch gate base (merge-base with target branch)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
@ -50,39 +56,47 @@ jobs:
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Clean Python cache
if: steps.changes.outputs.decision != 'skip'
run: |
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check uv.lock is up to date
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check || (echo "❌ uv.lock is out of sync with pyproject.toml. Run 'uv lock' locally and commit the result." && exit 1)
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
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
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Check ruff format
if: steps.changes.outputs.decision != 'skip'
run: |
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
if [ ! -s "$RUNNER_TEMP/ruff_format_files.txt" ]; then
@ -92,6 +106,7 @@ jobs:
xargs uv run --no-sync ruff format --check --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
- name: Debug - Check file state
if: steps.changes.outputs.decision != 'skip'
run: |
echo "Current branch:"
git branch --show-current
@ -101,30 +116,36 @@ jobs:
head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10
- name: Run Ruff linting
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync ruff check .
cd ..
- name: Check strict-rule budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/ruff_strict_gate.py --base "$GATE_BASE_SHA"
- name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA"
- name: Print OpenAI version
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- name: Check basedpyright budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
- name: Check tests/e2e basedpyright (zero errors)
if: steps.changes.outputs.decision != 'skip'
run: |
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
uv run --no-sync basedpyright tests/e2e
@ -133,12 +154,14 @@ jobs:
fi
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync python ../tests/documentation_tests/test_circular_imports.py
cd ..
- name: Check import safety
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)

View file

@ -1,6 +1,7 @@
name: UI Build Check
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -28,7 +29,14 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -36,7 +44,9 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Build
if: steps.changes.outputs.decision != 'skip'
run: npm run build

View file

@ -1,6 +1,7 @@
name: UI Unit Tests
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -32,7 +33,14 @@ jobs:
fetch-depth: 1
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -40,14 +48,17 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Run UI type tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
run: npm run test:types
- name: Run UI unit tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
GH_TOKEN: ${{ github.token }}

View file

@ -10,6 +10,7 @@ on:
permissions:
contents: read
pull-requests: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
@ -25,26 +26,34 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Thank You Message
run: |
echo "### 🙏 Thank you for contributing to LiteLLM!" >> $GITHUB_STEP_SUMMARY
echo "Your PR is being tested now. We appreciate your help in making LiteLLM better!" >> $GITHUB_STEP_SUMMARY
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check
.github/scripts/uv_sync_with_retries.sh --frozen --group proxy-dev --extra proxy --extra semantic-router
- name: Run MCP tests
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync pytest tests/mcp_tests -x -vv -n 4 --cov=./litellm --cov-report=xml --durations=5

View file

@ -32,28 +32,32 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
if: steps.changes.outputs.decision != 'skip'
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Detect backend-relevant changes
id: changes
uses: ./.github/actions/detect-backend-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |

View file

@ -221,7 +221,7 @@ overwrite_user_with_key_hash: bool = (
bedrock_request_metadata_fields: Optional[Sequence[str]] = (
None # allow-list of `user_api_key_*` fields (+ `spend_logs_metadata`) sent as Bedrock `requestMetadata`
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
store_audit_logs: bool | None = None
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False
### end of callbacks #############

View file

@ -1766,6 +1766,17 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
# one is seconds old, so a few minutes separates them.
PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300
# How long enqueued-token reservations for batches live without a refund. Providers
# complete or expire batches within their completion window (24h for OpenAI), so a
# reservation still unrefunded after 8 days belongs to a batch whose terminal state
# was never observed (e.g. proxy restart); expiry returns the tokens to the caller.
BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60
# Key/team metadata field that opts batches into enqueued-token limiting. Only proxy
# admins may write it: when present it replaces the standard RPM/TPM checks for
# batch submissions.
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit"
# Shared read-only empty mapping, for defaulting optional Mapping parameters without
# constructing a fresh mutable dict at each call site.
EMPTY_MAPPING: Final = MappingProxyType({})

View file

@ -0,0 +1,149 @@
"""Which deployments accrue PTU flat cost, and what that costs them per token.
Reserved provisioned throughput is billed by the hour whether or not requests are sent, so
a deployment that accrues flat cost must not also bill per token. The two halves live here
together because they have to agree: a deployment the rollup declines to charge but the
router prices at zero serves its traffic for free.
"""
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Final
from litellm.secret_managers.main import get_secret_bool
from litellm.types.router import ModelInfo
from litellm.types.utils import CustomPricingLiteLLMParams, MirroredPricingParams
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
def is_ptu_cost_attribution_enabled() -> bool:
"""Whether PTU flat-cost attribution is turned on for this process."""
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True
PTU_ZEROED_PRICING_FIELDS: Final = tuple(f for f in MirroredPricingParams.model_fields if f != "tiered_pricing") + (
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost_above_200k_tokens",
)
# tiered_pricing is emptied rather than zeroed: its tiers outrank the zeros written beside
# them, so a zero here would leave the cost map's tiers billing the traffic the reserved
# capacity already covers.
PTU_EMPTIED_PRICING_FIELDS: Final = frozenset(("tiered_pricing",))
# search_context_cost_per_query holds its rates in a table keyed by context size, and an
# absent table means the provider's own default rather than free, so it is zeroed in place
# and written on every PTU deployment rather than only where a table is already stored.
PTU_ZEROED_TABLE_FIELDS: Final = frozenset(("search_context_cost_per_query",))
SEARCH_CONTEXT_SIZES: Final = ("search_context_size_low", "search_context_size_medium", "search_context_size_high")
# Rate fields only. CustomPricingLiteLLMParams also carries settings that are not charges,
# and zeroing one of those would destroy the deployment's configuration rather than stop a
# charge.
CUSTOM_PRICING_FIELDS: Final = frozenset(f for f in CustomPricingLiteLLMParams.model_fields if "cost" in f)
PTU_ZEROED_PRICING: Final[Mapping[str, float | tuple[()] | Mapping[str, float]]] = MappingProxyType(
{
**dict.fromkeys(PTU_ZEROED_PRICING_FIELDS, 0.0),
**dict.fromkeys(PTU_EMPTIED_PRICING_FIELDS, ()),
**dict.fromkeys(PTU_ZEROED_TABLE_FIELDS, MappingProxyType(dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0))),
}
)
@dataclass(frozen=True, slots=True)
class PTUTerms:
"""The reservation a deployment declares, once every field has been validated."""
team_id: str
ptu_count: int
cost_per_ptu_per_hour: float
effective_from: datetime
effective_to: datetime | None
def _to_utc(parsed: datetime) -> datetime:
"""``parsed`` as UTC, reading a naive value as UTC rather than local time."""
return parsed.replace(tzinfo=timezone.utc) if parsed.tzinfo is None else parsed.astimezone(timezone.utc)
def _as_utc(value: object) -> datetime | None:
"""A model_info datetime as UTC, parsing an ISO string, else None."""
if isinstance(value, datetime):
return _to_utc(value)
if not isinstance(value, str):
return None
try:
return _to_utc(datetime.fromisoformat(value.replace("Z", "+00:00")))
except ValueError:
return None
def ptu_terms(model_info: Mapping[str, object]) -> PTUTerms | None:
"""The reservation this deployment accrues flat cost for, else None.
A start is required rather than inferred because flat cost accrues from it, and a
present but unparseable bound would read as no bound and widen the window to the whole
day, so either one leaves the deployment unpriced until the config is fixed.
"""
ptu_count: Final = model_info.get("ptu_count")
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
team_id: Final = model_info.get("team_id")
if ptu_count is None or cost_per_hour is None or not team_id:
return None
try:
ptu_count_int: Final = int(ptu_count)
cost_per_hour_float: Final = float(cost_per_hour)
except (TypeError, ValueError, OverflowError):
return None
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
return None
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
return None
raw_from: Final = model_info.get("ptu_effective_from")
raw_to: Final = model_info.get("ptu_effective_to")
effective_from: Final = _as_utc(raw_from)
effective_to: Final = _as_utc(raw_to)
if effective_from is None or (raw_to is not None and effective_to is None):
return None
if effective_to is not None and effective_to <= effective_from:
return None
return PTUTerms(
team_id=str(team_id),
ptu_count=ptu_count_int,
cost_per_ptu_per_hour=cost_per_hour_float,
effective_from=effective_from,
effective_to=effective_to,
)
def zeroed_ptu_pricing(
model_info: Mapping[str, object], declared: Mapping[str, object]
) -> Mapping[str, float | tuple[()] | Mapping[str, float]] | None:
"""The pricing a deployment accruing flat cost must carry, else None.
Both conditions hold or nothing is zeroed. Without the flag no flat cost accrues, so
zeroing would leave the deployment serving for free with nothing charged in its place,
which is what an SDK user who happens to carry ptu_count would otherwise get. The terms
are checked first only because they are a few dict reads, while the flag can resolve
through a configured secret manager, and this runs for every deployment registered.
Any further rate the deployment itself declares is zeroed alongside the standing set,
since one left standing bills the traffic the reserved capacity already paid for.
"""
if ptu_terms(model_info) is None:
return None
if not is_ptu_cost_attribution_enabled():
return None
return MappingProxyType(
{
**PTU_ZEROED_PRICING,
**dict.fromkeys(
CUSTOM_PRICING_FIELDS.intersection(declared)
.difference(PTU_ZEROED_TABLE_FIELDS)
.difference(PTU_EMPTIED_PRICING_FIELDS),
0.0,
),
}
)

View file

@ -10,6 +10,7 @@ from typing_extensions import TypedDict
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
@ -134,8 +135,11 @@ class BaseAnthropicMessagesStreamingIterator:
if self.completion_start_time is not None:
self.litellm_logging_obj.completion_start_time = self.completion_start_time
self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
# Enqueue on the rooted logging worker rather than asyncio.create_task:
# this also runs during generator teardown after a client disconnect,
# where an unrooted task could be garbage-collected before it bills.
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=self.litellm_logging_obj,
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
url_route="/v1/messages",
@ -197,13 +201,21 @@ class BaseAnthropicMessagesStreamingIterator:
collected_chunks: Final = []
saw_terminal_event = False
async for chunk in completion_stream:
if self.completion_start_time is None:
self.completion_start_time = datetime.now()
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
collected_chunks.append(encoded_chunk)
yield encoded_chunk
try:
async for chunk in completion_stream:
if self.completion_start_time is None:
self.completion_start_time = datetime.now()
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
collected_chunks.append(encoded_chunk)
yield encoded_chunk
except (GeneratorExit, asyncio.CancelledError):
# A client disconnect tears the generator down at the yield, so the
# post-loop logging below never runs and the tokens already streamed
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
if collected_chunks:
await self._handle_streaming_logging(collected_chunks)
raise
if not saw_terminal_event:
yield _incomplete_stream_error_sse_event()

View file

@ -842,8 +842,15 @@ class AmazonAnthropicClaudeMessagesConfig(
patched_stream: Final = self._promote_message_stop_usage(completion_stream)
async for chunk in handler.async_sse_wrapper(patched_stream):
yield chunk
sse_stream: Final = handler.async_sse_wrapper(patched_stream)
try:
async for chunk in sse_stream:
yield chunk
finally:
# Close the inner generator deterministically so a client disconnect
# (GeneratorExit here) reaches async_sse_wrapper's partial-spend logging
# now instead of at garbage collection. See LIT-5839.
await sse_stream.aclose()
@staticmethod
def _merge_message_start_cache_into_delta_usage(

View file

@ -12,7 +12,12 @@ from pydantic import PositiveInt, TypeAdapter, ValidationError
import litellm
from litellm import Router, provider_list
from litellm._logging import verbose_proxy_logger
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS
from litellm.constants import (
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
EMPTY_MAPPING,
MINIMUM_CUSTOM_KEY_LENGTH,
STANDARD_CUSTOMER_ID_HEADERS,
)
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import (
SSRFError,
@ -1169,6 +1174,46 @@ def enforce_output_token_estimates_are_admin_only(
)
class BatchEnqueuedTokenLimitRequest(Protocol):
"""The shape of any management request that can carry a batch enqueued-token limit."""
@property
def metadata(self) -> Mapping[str, object] | None: ...
@property
def model_fields_set(self) -> Collection[str]: ...
def enforce_batch_enqueued_token_limit_is_admin_only(
data: BatchEnqueuedTokenLimitRequest,
existing_metadata: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
entity: Literal["key", "team"],
) -> None:
"""Only a proxy admin may change a key or team's batch enqueued-token limit.
When set, ``batch_enqueued_token_limit`` replaces the standard RPM/TPM checks
for batch submissions, so a holder-writable copy would let a caller lift their
own batch quota. Gated on the resulting value rather than on presence, so a
form resending the stored value stays a no-op.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
stored: Final[Mapping[str, object]] = existing_metadata or EMPTY_MAPPING
requested: Final[Mapping[str, object]] = (
(data.metadata or EMPTY_MAPPING) if "metadata" in data.model_fields_set else stored
)
if requested.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) == stored.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY):
return
raise HTTPException(
status_code=403,
detail={ # mutable-ok: HTTPException.detail has no immutable form
"error": f"Only proxy admins can set {BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY} on a {entity}. "
"It replaces the standard rate limit checks for batch submissions."
},
)
def get_model_rate_limit_from_metadata(
user_api_key_dict: UserAPIKeyAuth,
metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"],

View file

@ -0,0 +1,456 @@
"""
Enqueued-token accounting for batch submissions.
Opt-in via admin-set ``batch_enqueued_token_limit`` in key or team metadata: batch
submissions reserve their estimated token count against a long-lived
enqueued-token allowance instead of the per-minute rate-limit windows, and
the reservation is refunded when the batch reaches a terminal state
(completed, failed, expired, or cancelled).
"""
import asyncio
import math
import time
import uuid
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.constants import BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, BATCH_ENQUEUED_TOKEN_TTL_SECONDS
from litellm.proxy._types import UserAPIKeyAuth
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
Span = _Span
InternalUsageCache = _InternalUsageCache
BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset(
{"completed", "complete", "failed", "expired", "cancelled", "cancelling"}
)
ScopeKey: TypeAlias = Literal["api_key", "team"]
RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """
local amount = tonumber(ARGV[1])
local ttl = tonumber(ARGV[2])
local limit = tonumber(ARGV[3])
local current = tonumber(redis.call('GET', KEYS[1]) or '0')
if current + amount > limit then
return {0, current}
end
local updated = redis.call('INCRBY', KEYS[1], amount)
redis.call('EXPIRE', KEYS[1], ttl)
return {1, updated}
"""
REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """
local updated = redis.call('DECRBY', KEYS[1], tonumber(ARGV[1]))
if updated <= 0 then
redis.call('DEL', KEYS[1])
end
return 1
"""
SAVE_RESERVATION_SCRIPT: Final = """
redis.call('SET', KEYS[1], ARGV[1], 'EX', tonumber(ARGV[2]))
return 1
"""
POP_RESERVATION_SCRIPT: Final = """
local value = redis.call('GET', KEYS[1])
if value and value ~= '' then
redis.call('SET', KEYS[1], '', 'EX', tonumber(ARGV[1]))
end
return value
"""
@dataclass(frozen=True, slots=True)
class BatchEnqueuedTokenScope:
key: ScopeKey
value: str
limit: int
ReservationBackend: TypeAlias = Literal["redis", "memory"]
@dataclass(frozen=True, slots=True)
class BatchEnqueuedTokenReservation:
tokens: int
scopes: tuple[BatchEnqueuedTokenScope, ...]
backend: ReservationBackend = "redis"
owner: str = ""
reserved_at_monotonic: float = field(default_factory=time.monotonic, compare=False)
@dataclass(frozen=True, slots=True)
class BatchEnqueuedTokenOverLimit:
scope: BatchEnqueuedTokenScope
enqueued: int
BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit
_LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)])
_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int])
_POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None)
_STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None)
_RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation)
class _ScriptRunner(Protocol):
def __call__(self, keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> Awaitable[object]: ...
def _read_metadata_limit(metadata: Mapping[str, object] | None) -> int | None:
if not metadata:
return None
raw: Final = metadata.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY)
if raw is None:
return None
try:
return _LIMIT_ADAPTER.validate_python(raw)
except ValidationError:
verbose_proxy_logger.warning(
"Ignoring invalid %s value %r; expected a positive integer",
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
raw,
)
return None
def resolve_batch_enqueued_token_scopes(
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[BatchEnqueuedTokenScope, ...]:
key_limit: Final = _read_metadata_limit(user_api_key_dict.metadata)
team_limit: Final = _read_metadata_limit(user_api_key_dict.team_metadata)
candidates: Final = (
BatchEnqueuedTokenScope(key="api_key", value=user_api_key_dict.api_key, limit=key_limit)
if key_limit is not None and user_api_key_dict.api_key
else None,
BatchEnqueuedTokenScope(key="team", value=user_api_key_dict.team_id, limit=team_limit)
if team_limit is not None and user_api_key_dict.team_id
else None,
)
return tuple(scope for scope in candidates if scope is not None)
def canonical_provider_batch_id(batch_id: str) -> str:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper
get_batch_id_from_unified_batch_id,
get_original_file_id,
)
decoded: Final = _is_base64_encoded_unified_file_id(batch_id)
if isinstance(decoded, str):
if "llm_batch_id" in decoded or "generic_response_id" in decoded:
return get_batch_id_from_unified_batch_id(decoded)
return decoded
return get_original_file_id(batch_id)
class _BatchResponseView(BaseModel):
model_config = ConfigDict(extra="ignore")
id: str
status: str
object: Literal["batch"]
def batch_response_view(response: object) -> _BatchResponseView | None:
try:
return _BatchResponseView.model_validate(response, from_attributes=True)
except ValidationError:
return None
class BatchEnqueuedTokenStore:
"""Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds.
Counters and records live in Redis when Redis is configured, through
single-key Lua scripts issued one scope at a time (Redis Cluster safe: no
cross-slot commands), with an over-limit or failing scope rolling back the
scopes reserved before it; otherwise a single-process in-memory fallback
guarded by one asyncio lock is used. Reservations remember which backend
granted them, and in-memory grants also remember the granting worker, so a
refund never debits counters the grant did not charge. Everything expires after
``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the
terminal-state refund can never leak tokens forever, and reservation records
expire no later than the counters they would refund, so a stale record can
never debit an allowance re-granted after its counters expired.
"""
def __init__(
self,
internal_usage_cache: "InternalUsageCache",
monotonic: Callable[[], float] = time.monotonic,
) -> None:
self.internal_usage_cache = internal_usage_cache
self._monotonic: Final = monotonic
self._lock = asyncio.Lock()
self._owner_token = uuid.uuid4().hex
redis_cache = internal_usage_cache.dual_cache.redis_cache
self._reserve_script: _ScriptRunner | None = (
redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None
)
self._refund_script: _ScriptRunner | None = (
redis_cache.async_register_script(REFUND_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None
)
self._save_script: _ScriptRunner | None = (
redis_cache.async_register_script(SAVE_RESERVATION_SCRIPT) if redis_cache is not None else None
)
self._pop_script: _ScriptRunner | None = (
redis_cache.async_register_script(POP_RESERVATION_SCRIPT) if redis_cache is not None else None
)
@staticmethod
def _counter_key(scope: BatchEnqueuedTokenScope) -> str:
return f"batch_enqueued_tokens:{scope.key}:{scope.value}"
@staticmethod
def _record_key(batch_id: str) -> str:
return f"batch_enqueued_token_reservation:{batch_id}"
async def reserve(
self,
tokens: int,
scopes: tuple[BatchEnqueuedTokenScope, ...],
litellm_parent_otel_span: "Span | None" = None,
) -> BatchEnqueuedTokenOutcome:
if tokens <= 0 or not scopes:
return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes)
reserve_script: Final = self._reserve_script
refund_script: Final = self._refund_script
if reserve_script is not None and refund_script is not None:
try:
return await self._reserve_via_redis(reserve_script, refund_script, tokens=tokens, scopes=scopes)
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters
verbose_proxy_logger.warning(
"Redis enqueued-token reserve failed, falling back to in-memory: %s", str(e)
)
return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span)
async def _reserve_via_redis(
self,
reserve_script: _ScriptRunner,
refund_script: _ScriptRunner,
tokens: int,
scopes: tuple[BatchEnqueuedTokenScope, ...],
) -> BatchEnqueuedTokenOutcome:
started: Final = self._monotonic()
for index, scope in enumerate(scopes):
result = await self._run_reserve_script(
reserve_script,
refund_script,
tokens=tokens,
scope=scope,
already_reserved=scopes[:index],
)
if result[0] != 1:
await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=scopes[:index])
return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1])
return BatchEnqueuedTokenReservation(
tokens=tokens, scopes=scopes, backend="redis", reserved_at_monotonic=started
)
async def _run_reserve_script(
self,
reserve_script: _ScriptRunner,
refund_script: _ScriptRunner,
tokens: int,
scope: BatchEnqueuedTokenScope,
already_reserved: tuple[BatchEnqueuedTokenScope, ...],
) -> tuple[int, int]:
try:
raw_result: Final = await reserve_script(
(self._counter_key(scope),),
(tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit),
)
return _RESERVE_RESULT_ADAPTER.validate_python(raw_result)
except Exception:
await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=already_reserved)
raise
async def _rollback_partial_reserve(
self,
refund_script: _ScriptRunner,
tokens: int,
scopes: tuple[BatchEnqueuedTokenScope, ...],
) -> None:
try:
await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes)
except Exception as e: # noqa: BLE001 # best-effort rollback: the leak is TTL-bounded and only tightens the allowance
verbose_proxy_logger.warning(
"Rollback of partially reserved enqueued tokens failed; leaked increments expire with the TTL: %s",
str(e),
)
async def _refund_via_redis(
self,
refund_script: _ScriptRunner,
tokens: int,
scopes: tuple[BatchEnqueuedTokenScope, ...],
) -> None:
for scope in scopes:
await refund_script((self._counter_key(scope),), (tokens,))
async def _reserve_in_memory(
self,
tokens: int,
scopes: tuple[BatchEnqueuedTokenScope, ...],
span: "Span | None",
) -> BatchEnqueuedTokenOutcome:
started: Final = self._monotonic()
async with self._lock:
currents: Final = tuple([await self._get_local_counter(scope, span) for scope in scopes])
for scope, current in zip(scopes, currents):
if current + tokens > scope.limit:
return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current)
for scope, current in zip(scopes, currents):
await self._set_local_counter(scope, current + tokens, span)
return BatchEnqueuedTokenReservation(
tokens=tokens, scopes=scopes, backend="memory", owner=self._owner_token, reserved_at_monotonic=started
)
async def refund(
self,
reservation: BatchEnqueuedTokenReservation,
litellm_parent_otel_span: "Span | None" = None,
) -> None:
if reservation.tokens <= 0 or not reservation.scopes:
return
if reservation.backend == "redis":
await self._refund_redis_reservation(reservation)
return
if reservation.owner != self._owner_token:
verbose_proxy_logger.warning(
"Skipping enqueued-token refund granted in another worker's memory; its counters expire with the TTL"
)
return
async with self._lock:
for scope in reservation.scopes:
current = await self._get_local_counter(scope, litellm_parent_otel_span)
remaining = current - reservation.tokens
if remaining <= 0:
self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._counter_key(scope))
else:
await self._set_local_counter(scope, remaining, litellm_parent_otel_span)
async def _refund_redis_reservation(self, reservation: BatchEnqueuedTokenReservation) -> None:
refund_script: Final = self._refund_script
if refund_script is None:
verbose_proxy_logger.warning(
"No Redis client for a Redis-granted enqueued-token refund; leaked increments expire with the TTL"
)
return
try:
await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes)
except Exception as e: # noqa: BLE001 # best-effort refund: the leak is TTL-bounded and only tightens the allowance
verbose_proxy_logger.warning(
"Redis enqueued-token refund failed; leaked increments expire with the TTL: %s", str(e)
)
async def save_reservation(
self,
batch_id: str,
reservation: BatchEnqueuedTokenReservation,
litellm_parent_otel_span: "Span | None" = None,
) -> None:
serialized: Final = _RESERVATION_ADAPTER.dump_json(reservation).decode("utf-8")
elapsed: Final = self._monotonic() - reservation.reserved_at_monotonic
ttl: Final = max(1, BATCH_ENQUEUED_TOKEN_TTL_SECONDS - math.ceil(elapsed))
if self._save_script is not None:
try:
await self._save_script(
(self._record_key(batch_id),),
(serialized, ttl),
)
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record
verbose_proxy_logger.warning(
"Redis enqueued-token reservation save failed, falling back to in-memory: %s", str(e)
)
else:
return
await self.internal_usage_cache.async_set_cache(
key=self._record_key(batch_id),
value=serialized,
ttl=ttl,
litellm_parent_otel_span=litellm_parent_otel_span,
local_only=True,
)
async def pop_reservation(
self,
batch_id: str,
litellm_parent_otel_span: "Span | None" = None,
) -> BatchEnqueuedTokenReservation | None:
redis_raw: Final = await self._pop_redis_record(batch_id)
if redis_raw is not None and not redis_raw:
# The Redis pop tombstones popped records in place, so a hit on the empty
# tombstone means the batch was already refunded elsewhere; a local copy
# left behind by a save that raised after landing must not refund again.
await self._pop_local_record(batch_id, litellm_parent_otel_span)
return None
raw: Final = (
redis_raw if redis_raw is not None else await self._pop_local_record(batch_id, litellm_parent_otel_span)
)
if raw is None:
return None
try:
if isinstance(raw, (str, bytes)):
return _RESERVATION_ADAPTER.validate_json(raw)
return _RESERVATION_ADAPTER.validate_python(raw)
except ValidationError:
verbose_proxy_logger.warning("Discarding malformed enqueued-token reservation record for %s", batch_id)
return None
async def _pop_redis_record(self, batch_id: str) -> str | bytes | None:
pop_script: Final = self._pop_script
if pop_script is None:
return None
try:
return _POPPED_VALUE_ADAPTER.validate_python(
await pop_script((self._record_key(batch_id),), (BATCH_ENQUEUED_TOKEN_TTL_SECONDS,))
)
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record
verbose_proxy_logger.warning(
"Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e)
)
return None
async def _pop_local_record(self, batch_id: str, span: "Span | None") -> object:
async with self._lock:
stored = await self.internal_usage_cache.async_get_cache(
key=self._record_key(batch_id),
litellm_parent_otel_span=span,
local_only=True,
)
if stored is None:
return None
self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._record_key(batch_id))
return stored
async def _get_local_counter(self, scope: BatchEnqueuedTokenScope, span: "Span | None") -> int:
stored = await self.internal_usage_cache.async_get_cache(
key=self._counter_key(scope),
litellm_parent_otel_span=span,
local_only=True,
)
return _STORED_COUNTER_ADAPTER.validate_python(stored) or 0
async def _set_local_counter(self, scope: BatchEnqueuedTokenScope, value: int, span: "Span | None") -> None:
await self.internal_usage_cache.async_set_cache(
key=self._counter_key(scope),
value=value,
ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS,
litellm_parent_otel_span=span,
local_only=True,
)

View file

@ -46,9 +46,16 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
map_v3_rate_limit_type,
)
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenOverLimit,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
resolve_batch_enqueued_token_scopes,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
PROJECT_ITPM_DESCRIPTOR_KEY,
PROJECT_OTPM_DESCRIPTOR_KEY,
get_or_create_request_stash,
)
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
@ -291,6 +298,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
has_enqueued_scopes: bool = False,
) -> tuple[bool, list["RateLimitDescriptor"] | None]:
"""
Skip downloading batch input files when the operator disabled batch
@ -343,8 +351,10 @@ class _PROXY_BatchRateLimiter(CustomLogger):
user_api_key_dict=user_api_key_dict,
data=data,
)
if not self._has_applicable_batch_rate_limits(descriptors) and not self._project_has_any_io_token_limits(
user_api_key_dict
if (
not has_enqueued_scopes
and not self._has_applicable_batch_rate_limits(descriptors)
and not self._project_has_any_io_token_limits(user_api_key_dict)
):
verbose_proxy_logger.debug("Skipping batch input file processing: no rate limits configured")
return True, None
@ -511,6 +521,59 @@ class _PROXY_BatchRateLimiter(CustomLogger):
return file_id, fetch_kwargs
async def _reserve_batch_enqueued_tokens(
self,
user_api_key_dict: UserAPIKeyAuth,
data: Mapping[str, object],
batch_usage: BatchFileUsage,
scopes: tuple[BatchEnqueuedTokenScope, ...],
) -> None:
"""Reserve the batch's estimated tokens against the caller's enqueued-token allowance.
Runs instead of the per-minute counter charge when the key or team
opted in via ``batch_enqueued_token_limit`` metadata. The reservation
is stashed on the request so the v3 limiter's post-call hooks can
persist it (keyed by the provider batch id) and refund it when the
batch reaches a terminal state.
"""
outcome: Final = await self.parallel_request_limiter.batch_enqueued_token_store.reserve(
tokens=batch_usage.total_tokens,
scopes=scopes,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
match outcome:
case BatchEnqueuedTokenOverLimit():
self._raise_enqueued_limit_error(over_limit=outcome, data=data, batch_usage=batch_usage)
case BatchEnqueuedTokenReservation():
get_or_create_request_stash().batch_enqueued_reservation = outcome
def _raise_enqueued_limit_error(
self,
over_limit: BatchEnqueuedTokenOverLimit,
data: Mapping[str, object],
batch_usage: BatchFileUsage,
) -> NoReturn:
scope: Final = over_limit.scope
remaining: Final = max(0, scope.limit - over_limit.enqueued)
detail: Final = (
f"Batch enqueued token limit exceeded for {scope.key}: {scope.value}. "
f"Batch requires {batch_usage.total_tokens} tokens but only {remaining} enqueued tokens remaining "
f"out of {scope.limit} enqueued token limit. "
f"Tokens free up as running batches complete or are cancelled."
)
raw_model: Final = data.get("model")
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
raw_model if isinstance(raw_model, str) else None
)
raise ProxyRateLimitError(
detail=detail,
headers=MappingProxyType({"rate_limit_type": "tokens"}),
category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT,
rate_limit_type=map_v3_rate_limit_type("tokens"),
model=resolved_model,
llm_provider=llm_provider,
)
def _raise_rate_limit_error(
self,
status: "RateLimitStatus",
@ -1039,8 +1102,9 @@ class _PROXY_BatchRateLimiter(CustomLogger):
verbose_proxy_logger.debug("No input_file_id in batch request, skipping rate limiting")
return data
enqueued_scopes: Final = resolve_batch_enqueued_token_scopes(user_api_key_dict)
should_skip, batch_rate_limit_descriptors = self._should_skip_batch_input_file_processing(
data=data, user_api_key_dict=user_api_key_dict
data=data, user_api_key_dict=user_api_key_dict, has_enqueued_scopes=bool(enqueued_scopes)
)
if should_skip:
return data
@ -1066,6 +1130,16 @@ class _PROXY_BatchRateLimiter(CustomLogger):
data["_batch_token_count"] = batch_usage.total_tokens
data["_batch_request_count"] = batch_usage.request_count
if enqueued_scopes:
await self._reserve_batch_enqueued_tokens(
user_api_key_dict=user_api_key_dict,
data=data,
batch_usage=batch_usage,
scopes=enqueued_scopes,
)
verbose_proxy_logger.debug("Batch enqueued-token reservation succeeded")
return data
# Directly increment counters by batch amounts (check happens atomically)
# This will raise HTTPException if limits are exceeded
await self._check_and_increment_batch_counters(

View file

@ -43,6 +43,7 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -53,8 +54,7 @@ class KeyManagementEventHooks:
except Exception as e:
verbose_proxy_logger.warning("Failed to send key created email: %s", e)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values: Final = response.model_dump_json(exclude_none=True)
asyncio.create_task(
create_audit_log_for_update(
@ -103,11 +103,11 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values: Final = json.dumps(data.json(exclude_none=True), default=str)
_before_value = existing_key_row.json(exclude_none=True)
@ -144,6 +144,7 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -180,7 +181,7 @@ class KeyManagementEventHooks:
verbose_proxy_logger.warning("Failed to send key rotated email: %s", e)
# store the audit log
if litellm.store_audit_logs is True and existing_key_row.token is not None:
if is_audit_logging_enabled() and existing_key_row.token is not None:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -218,12 +219,12 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True and data.keys is not None:
if is_audit_logging_enabled() and data.keys is not None:
# make an audit log for each key deleted
for key in keys_being_deleted:
if key.token is None:

View file

@ -44,6 +44,13 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
map_v3_rate_limit_type,
)
from litellm.proxy.hooks.batch_enqueued_tokens import (
BATCH_ENQUEUED_REFUND_STATUSES,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenStore,
batch_response_view,
canonical_provider_batch_id,
)
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage
@ -515,6 +522,7 @@ class RequestRateLimiterStash:
otpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field(
default_factory=frozenset
)
batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None
reservation_released: bool = False
@ -619,6 +627,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Batch rate limiter (lazy loaded)
self._batch_rate_limiter: CallTypeRateLimiter | None = None
self.batch_enqueued_token_store = BatchEnqueuedTokenStore(internal_usage_cache=internal_usage_cache)
# Serializes multi-phase check+increment sequences (batch + dynamic
# limiters) within this process to close the TOCTOU window between
@ -4673,6 +4682,32 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
except Exception as e:
verbose_proxy_logger.exception("Error in rate limit post-call hook: %s", e)
try:
await self._handle_batch_enqueued_post_call(user_api_key_dict=user_api_key_dict, response=response)
except Exception as e: # noqa: BLE001 # post-call batch accounting must never fail the response
verbose_proxy_logger.exception("Error in batch enqueued-token post-call hook: %s", e)
async def _handle_batch_enqueued_post_call(self, user_api_key_dict: UserAPIKeyAuth, response: object) -> None:
view: Final = batch_response_view(response)
if view is None:
return
span: Final = user_api_key_dict.parent_otel_span
stash: Final = get_request_stash()
if stash is not None and stash.batch_enqueued_reservation is not None:
await self.batch_enqueued_token_store.save_reservation(
batch_id=canonical_provider_batch_id(view.id),
reservation=stash.batch_enqueued_reservation,
litellm_parent_otel_span=span,
)
stash.batch_enqueued_reservation = None
if view.status.lower() in BATCH_ENQUEUED_REFUND_STATUSES:
popped: Final = await self.batch_enqueued_token_store.pop_reservation(
batch_id=canonical_provider_batch_id(view.id),
litellm_parent_otel_span=span,
)
if popped is not None:
await self.batch_enqueued_token_store.refund(reservation=popped, litellm_parent_otel_span=span)
async def async_post_call_failure_hook(
self,
request_data: dict,
@ -4706,6 +4741,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
stash.parallel_slot = None
if stash.batch_enqueued_reservation is not None:
await self.batch_enqueued_token_store.refund(
reservation=stash.batch_enqueued_reservation,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
stash.batch_enqueued_reservation = None
if stash.reservation_released:
return
reserved_tokens: Final = stash.reserved_tokens

View file

@ -20,7 +20,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
WebhookEvent,
)
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.repositories.user_repository import UserRepository
@ -203,7 +206,7 @@ class UserManagementEventHooks:
- user_api_key_dict: UserAPIKeyAuth - The user api key dictionary.
- litellm_proxy_admin_name: Optional[str] - The name of the proxy admin.
"""
if not litellm.store_audit_logs:
if not is_audit_logging_enabled():
return
from litellm.proxy.management_helpers.audit_logs import (

View file

@ -17,7 +17,6 @@ from typing import TYPE_CHECKING, Any, Final, Protocol
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._redis import _redis_kwargs_from_environment
from litellm._uuid import uuid
@ -299,14 +298,15 @@ async def _emit_cache_settings_audit_log(
exception. Captured under ``LiteLLM_CacheConfig`` so the row
co-locates with the table it mutates.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -100,16 +100,15 @@ async def _emit_hashicorp_vault_audit_log(
``LiteLLM_ConfigOverrides`` so the row co-locates with the table it
mutates.
"""
import litellm
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -243,12 +243,15 @@ async def _emit_coordination_redis_audit_log(
litellm_changed_by: str | None,
) -> None:
"""Emit an audit-log row for a /coordination_redis/settings mutation."""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -2220,6 +2220,7 @@ async def delete_user(
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -2298,9 +2299,8 @@ async def delete_user(
},
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
# make an audit log for each team deleted
_user_row = user_row.json(exclude_none=True)

View file

@ -57,6 +57,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
enforce_batch_enqueued_token_limit_is_admin_only,
enforce_output_token_estimates_are_admin_only,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -901,6 +902,12 @@ async def _common_key_generation_helper(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_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(
@ -2302,6 +2309,14 @@ async def _process_single_key_update(
prisma_client=prisma_client,
)
_existing_row_metadata: Final = getattr(existing_key_row, "metadata", None)
enforce_batch_enqueued_token_limit_is_admin_only(
data=update_key_request,
existing_metadata=_existing_row_metadata if isinstance(_existing_row_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,
entity="key",
)
# Check team member permissions
if prisma_client is not None:
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
@ -2564,6 +2579,12 @@ async def _validate_update_key_data(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_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
@ -4760,6 +4781,12 @@ async def _execute_virtual_key_regeneration(
user_api_key_dict=user_api_key_dict,
entity="key",
)
enforce_batch_enqueued_token_limit_is_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)
@ -6245,6 +6272,7 @@ async def block_key(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -6291,7 +6319,7 @@ async def block_key(
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -6358,6 +6386,7 @@ async def unblock_key(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -6404,7 +6433,7 @@ async def unblock_key(
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -64,7 +64,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.management_helpers.audit_logs import get_audit_log_changed_by
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.repositories.table_repositories import (
MCPServerRepository,
MCPUserCredentialsRepository,
@ -2018,7 +2021,7 @@ if MCP_AVAILABLE:
await global_mcp_server_manager.reload_servers_from_database()
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:
if is_audit_logging_enabled():
pass
# TODO: Delete from virtual keys
@ -2613,7 +2616,7 @@ if MCP_AVAILABLE:
)
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:
if is_audit_logging_enabled():
pass
return _redact_mcp_credentials(mcp_server_record_updated)

View file

@ -24,6 +24,13 @@ from pydantic import BaseModel, ConfigDict, Field, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.litellm_core_utils.ptu_pricing import (
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
PTU_ZEROED_PRICING_FIELDS,
PTU_ZEROED_TABLE_FIELDS,
SEARCH_CONTEXT_SIZES,
)
from litellm.proxy._types import (
BlockModelRequest,
CommonProxyErrors,
@ -89,7 +96,6 @@ from litellm.types.router import (
ModelInfo,
updateDeployment,
)
from litellm.types.utils import CustomPricingLiteLLMParams
from litellm.utils import get_utc_datetime
router: Final = APIRouter()
@ -346,12 +352,8 @@ def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None:
# tiered_pricing is the one mirrored field that is a table of ranges, not a rate, so it is stored
# empty (see _PTU_EMPTIED_PRICING_FIELDS): its tiers outrank the zeros written beside them, so
# dropping it would leave the cost map's tiers billing the traffic the reserved capacity covers.
_PTU_ZEROED_PRICING_FIELDS: Final = tuple(f for f in SPECIAL_MODEL_INFO_PARAMS if f != "tiered_pricing") + (
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost_above_200k_tokens",
)
_PTU_EMPTIED_PRICING_FIELDS: Final = frozenset({"tiered_pricing"})
_PTU_ZEROED_PRICING_FIELDS: Final = PTU_ZEROED_PRICING_FIELDS
_PTU_EMPTIED_PRICING_FIELDS: Final = PTU_EMPTIED_PRICING_FIELDS
_PTU_ZEROED_PRICING: Final[Mapping[str, float | tuple[()]]] = MappingProxyType(
{
**dict.fromkeys(_PTU_ZEROED_PRICING_FIELDS, 0.0),
@ -363,13 +365,13 @@ _EMPTY_MODEL_INFO: Final[Mapping[str, object]] = _NO_PRICING_OVERRIDE
# Rate fields only. CustomPricingLiteLLMParams also carries settings that are not charges
# (an embedding's output_vector_size, the regional uplift multipliers), and zeroing one of
# those would destroy the deployment's configuration rather than stop a charge.
_CUSTOM_PRICING_FIELDS: Final = frozenset(f for f in CustomPricingLiteLLMParams.model_fields if "cost" in f)
_CUSTOM_PRICING_FIELDS: Final = CUSTOM_PRICING_FIELDS
# search_context_cost_per_query holds its rates in a table keyed by context size, and an absent
# table means the provider's own default rate rather than free (litellm/llms/gemini/cost_calculator
# falls back to $0.035), so it is zeroed in place rather than emptied like tiered_pricing, and
# written on every PTU deployment rather than only where a table is already stored.
_PTU_ZEROED_TABLE_FIELDS: Final = frozenset({"search_context_cost_per_query"})
_SEARCH_CONTEXT_SIZES: Final = ("search_context_size_low", "search_context_size_medium", "search_context_size_high")
_PTU_ZEROED_TABLE_FIELDS: Final = PTU_ZEROED_TABLE_FIELDS
_SEARCH_CONTEXT_SIZES: Final = SEARCH_CONTEXT_SIZES
def _is_nonzero_rate(value: object) -> bool:

View file

@ -13,7 +13,6 @@ from typing import Annotated, Any, Final
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import (
@ -182,14 +181,15 @@ async def _emit_team_callback_audit_log(
Callback secrets are redacted before serialization so the audit table
cannot itself become a credential-harvest sink.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
redacted_before: Final = _redact_callback_secrets(before_metadata)
redacted_after: Final = _redact_callback_secrets(after_metadata)

View file

@ -85,7 +85,10 @@ 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.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
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
@ -1249,6 +1252,7 @@ async def new_team(
try:
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
_license_check,
@ -1303,6 +1307,12 @@ async def new_team(
user_api_key_dict=user_api_key_dict,
entity="team",
)
enforce_batch_enqueued_token_limit_is_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()
@ -1551,8 +1561,7 @@ async def new_team(
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values = complete_team_data.json(exclude_none=True)
_updated_values = json.dumps(_updated_values, default=str)
@ -1944,6 +1953,7 @@ async def update_team(
```
"""
try:
from litellm.proxy.management_helpers.audit_logs import is_audit_logging_enabled
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
@ -2007,6 +2017,12 @@ async def update_team(
user_api_key_dict=user_api_key_dict,
entity="team",
)
enforce_batch_enqueued_token_limit_is_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")
@ -2246,8 +2262,7 @@ async def update_team(
proxy_logging_obj=proxy_logging_obj,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
await _create_team_update_audit_log(
existing_team_row=existing_team_row,
updated_kv=updated_kv,
@ -3712,6 +3727,7 @@ async def delete_team(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -3756,9 +3772,8 @@ async def delete_team(
litellm_changed_by=litellm_changed_by,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
# make an audit log for each team deleted
for team_id in data.team_ids:
team_row: LiteLLM_TeamTable | None = await prisma_client.get_data(

View file

@ -24,6 +24,22 @@ _audit_log_callback_cache: Final[dict[str, CustomLogger]] = {}
ALLOW_LITELLM_CHANGED_BY_HEADER_METADATA_KEY: Final = "allow_litellm_changed_by_header"
def is_audit_logging_enabled(store_audit_logs: bool | None = None) -> bool:
from litellm.secret_managers.main import get_secret_bool
configured_value: Final[bool | None] = litellm.store_audit_logs if store_audit_logs is None else store_audit_logs
if configured_value is not None:
return configured_value
environment_value: Final[bool | None] = get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if environment_value is not None:
return environment_value
from litellm.proxy.proxy_server import premium_user
return premium_user is True
def _allows_litellm_changed_by_header(user_api_key_dict: UserAPIKeyAuth) -> bool:
for admin_metadata in (user_api_key_dict.metadata, user_api_key_dict.team_metadata):
if (
@ -164,11 +180,7 @@ async def create_object_audit_log(
- user_api_key_dict: UserAPIKeyAuth - The user api key dictionary.
- litellm_proxy_admin_name: Optional[str] - The name of the proxy admin.
"""
from litellm.secret_managers.main import get_secret_bool
_store_audit_logs: Final[bool | None] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if _store_audit_logs is not True:
if not is_audit_logging_enabled():
return
_changed_by: Final = get_audit_log_changed_by(
@ -196,10 +208,7 @@ async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs):
"""
Create an audit log for an object.
"""
from litellm.secret_managers.main import get_secret_bool
_store_audit_logs: Final[bool | None] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if _store_audit_logs is not True:
if not is_audit_logging_enabled():
return
from litellm.proxy.proxy_server import premium_user, prisma_client

View file

@ -5035,6 +5035,7 @@ class ProxyConfig:
)
elif key == "audit_log_callbacks":
from litellm.proxy.management_helpers.audit_logs import (
is_audit_logging_enabled,
reset_audit_log_callback_cache,
)
@ -5053,14 +5054,14 @@ class ProxyConfig:
litellm.audit_log_callbacks.append(callback)
_store_audit_logs = litellm_settings.get("store_audit_logs", litellm.store_audit_logs)
if _store_audit_logs:
if is_audit_logging_enabled(store_audit_logs=_store_audit_logs):
print( # noqa: T201
f"{blue_color_code} Initialized Audit Log Callbacks - {litellm.audit_log_callbacks} {reset_color_code}"
)
else:
verbose_proxy_logger.warning(
"'audit_log_callbacks' is configured but 'store_audit_logs' is not enabled. "
"Audit log callbacks will not fire until 'store_audit_logs: true' is added to litellm_settings."
"'audit_log_callbacks' is configured but audit logging is not enabled. "
"Audit log callbacks will not fire."
)
elif key == "cache_params":
# this is set in the cache branch

View file

@ -1,18 +1,12 @@
"""Opt-in flag for PTU (provisioned throughput unit) flat-cost attribution.
"""Re-exported from ``litellm.litellm_core_utils.ptu_pricing``.
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.
The flag lives in core because the router reads it while registering a deployment, and
router code cannot import from the proxy.
"""
from typing import Final
from litellm.litellm_core_utils.ptu_pricing import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
)
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
__all__ = ("PTU_COST_ATTRIBUTION_ENV_VAR", "is_ptu_cost_attribution_enabled")

View file

@ -14,6 +14,7 @@ and share the existing unique constraint.
import asyncio
import json
import sys
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from datetime import date, datetime, time, timedelta, timezone
@ -29,14 +30,15 @@ from litellm.constants import (
PTU_ROLLUP_MAX_BACKFILL_DAYS,
PTU_SENTINEL_API_KEY,
)
from litellm.litellm_core_utils.ptu_pricing import ptu_terms
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.types.router import ModelInfo
if TYPE_CHECKING:
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.utils import PrismaClient
_HOURS_PER_DAY: Final = 24
_PRUNE_ID_CHUNK_SIZE: Final = 5_000
_UPSERT_ATTEMPTS: Final = 3
_UPSERT_RETRY_BACKOFF_SECONDS: Final = 0.5
@ -72,28 +74,6 @@ class PTUModel:
effective_to: datetime | None = None
def _parse_utc_datetime(value: object) -> datetime | None:
"""Parse a model_info datetime (ISO string or datetime) into a UTC-aware datetime, else None."""
parsed: Final = _coerce_datetime(value)
if parsed is None:
return None
if parsed.tzinfo is None:
return parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
def _coerce_datetime(value: object) -> datetime | None:
"""``value`` as a datetime, parsing an ISO string, else None."""
if isinstance(value, datetime):
return value
if not isinstance(value, str):
return None
try:
return datetime.fromisoformat(value.replace("Z", "+00:00"))
except ValueError:
return None
def _public_model_name(row: object, model_info: Mapping[str, object]) -> str:
"""The name an operator recognises for this deployment.
@ -167,46 +147,20 @@ def _parse_ptu_model(row: object) -> PTUModel | None:
Valid means model_info has a positive ptu_count, a non-negative
cost_per_ptu_per_hour, and a team_id (1 model -> 1 team).
"""
raw_model_info: Final = getattr(row, "model_info", None)
model_info: Final = _decode_model_info(raw_model_info)
model_info: Final = _decode_model_info(getattr(row, "model_info", None))
if model_info is None:
return None
ptu_count: Final = model_info.get("ptu_count")
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
team_id: Final = model_info.get("team_id")
if ptu_count is None or cost_per_hour is None or not team_id:
return None
try:
ptu_count_int: Final = int(ptu_count)
cost_per_hour_float: Final = float(cost_per_hour)
except (TypeError, ValueError, OverflowError):
return None
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
return None
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
return None
if model_info.get("ptu_effective_from") is None:
# The endpoints require a start; a row without one predates that rule or was
# written around them, and inferring one would bill days the deployment did not exist
return None
raw_from: Final = model_info.get("ptu_effective_from")
raw_to: Final = model_info.get("ptu_effective_to")
effective_from: Final = _parse_utc_datetime(raw_from)
effective_to: Final = _parse_utc_datetime(raw_to)
# A present-but-unparseable bound would read as "no bound" and silently widen the
# window to the whole day, so the deployment is skipped until the config is fixed
if (raw_from is not None and effective_from is None) or (raw_to is not None and effective_to is None):
return None
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
terms: Final = ptu_terms(model_info)
if terms is None:
return None
return PTUModel(
model_id=str(getattr(row, "model_id", "") or ""),
model_name=_public_model_name(row, model_info),
team_id=str(team_id),
ptu_count=ptu_count_int,
cost_per_ptu_per_hour=cost_per_hour_float,
effective_from=effective_from,
effective_to=effective_to,
team_id=terms.team_id,
ptu_count=terms.ptu_count,
cost_per_ptu_per_hour=terms.cost_per_ptu_per_hour,
effective_from=terms.effective_from,
effective_to=terms.effective_to,
)
@ -358,10 +312,70 @@ async def _upsert_charge_with_retry(
return False
async def _load_ptu_models(prisma_client: "PrismaClient") -> tuple[PTUModel, ...]:
"""Every model deployment currently carrying valid manual PTU config."""
@dataclass(frozen=True, slots=True)
class _LoadedDeployments:
"""The deployments a run will price, and every deployment id it looked at.
The id set is deliberately wider than the priced set. A deployment whose PTU config
was removed produces no charge and still has to be prunable, so bounding the prune on
what priced would strand its old rows forever. It is also a guaranteed superset of the
priced set, or a run could write a charge that falls outside its own delete filter.
"""
models: tuple[PTUModel, ...]
scanned_ids: frozenset[str]
config_sourced: bool
def _running_router() -> object | None:
"""The proxy's router, or None outside a running proxy.
Read out of ``sys.modules`` rather than imported, so a rollup driven from a test or a
script does not pull the whole proxy server in behind it.
"""
proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server")
return getattr(proxy_server, "llm_router", None) if proxy_server is not None else None
def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -> tuple[_PTUDeployment, ...]:
"""Deployments the router holds that no ``LiteLLM_ProxyModelTable`` row owns.
``db_model`` is forced True on every deployment loaded from that table and defaults to
False on ModelInfo, so the complement is what config.yaml declared. A per-request
credential clone carries ``original_model_id`` and reuses its source's PTU config under
a fresh id, so pricing it would bill one reservation once per distinct client key.
"""
entries: Final = tuple(getattr(router, "model_list", None) or ())
records: Final = tuple(_router_deployment(entry) for entry in entries)
return tuple(
record
for record in records
if record is not None
and record.model_info.get("db_model") is not True
and record.model_info.get("original_model_id") is None
and record.model_id not in owned_by_db
)
async def _load_ptu_models(prisma_client: "PrismaClient") -> _LoadedDeployments:
"""Every deployment carrying valid manual PTU config, and every id the scan saw.
Reserved capacity is billed by the provider whichever file declared it, so a
deployment the proxy only knows from config.yaml accrues alongside the stored ones.
"""
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
return tuple(parsed for parsed in (_parse_ptu_model(row) for row in rows) if parsed is not None)
db_ids: Final = frozenset(model_id for row in rows if (model_id := str(getattr(row, "model_id", "") or "")))
config_records: Final = _config_deployments(_running_router(), owned_by_db=db_ids)
models: Final = tuple(
parsed for parsed in (_parse_ptu_model(row) for row in (*rows, *config_records)) if parsed is not None
)
return _LoadedDeployments(
models=models,
config_sourced=bool(config_records),
scanned_ids=db_ids
| frozenset(record.model_id for record in config_records)
| frozenset(model.model_id for model in models),
)
async def run_ptu_flat_cost_rollup(
@ -378,8 +392,10 @@ async def run_ptu_flat_cost_rollup(
The prune predicate is ``updated_at < run_started`` rather than "not in the charge
set I computed", which matters under concurrency: whether a row is garbage becomes a
property of the row instead of one run's in-memory config snapshot, so a run can
never delete a row a concurrent run just wrote. It is still skipped when any charge
failed to write, since a row whose replacement never landed would look unrefreshed.
never delete a row a concurrent run just wrote. It is bounded to the deployments this
run looked at, so a row it cannot account for is out of reach either way. It is still
skipped when any charge failed to write, since a row whose replacement never landed
would look unrefreshed.
"""
day: Final = target_date or (datetime.now(timezone.utc).date() - timedelta(days=1))
@ -390,7 +406,8 @@ async def run_ptu_flat_cost_rollup(
date_str: Final = day.isoformat()
run_started: Final = datetime.now(timezone.utc)
ptu_models: Final = await _load_ptu_models(prisma_client)
loaded: Final = await _load_ptu_models(prisma_client)
ptu_models: Final = loaded.models
charges: Final = _aggregate_charges(ptu_models, day)
landed: Final = tuple(
@ -415,7 +432,12 @@ async def run_ptu_flat_cost_rollup(
date_str,
)
else:
await _prune_unrefreshed_sentinel_rows(prisma_client, date_str=date_str, run_started=run_started)
await _prune_unrefreshed_sentinel_rows(
prisma_client,
date_str=date_str,
run_started=run_started,
scanned_ids=loaded.scanned_ids if loaded.config_sourced else None,
)
verbose_proxy_logger.info(
"PTU rollup for %s: %d PTU models processed, %d rows written, %d rows failed",
@ -524,7 +546,7 @@ async def run_ptu_flat_cost_backfill(
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
ptu_models: Final = await _load_ptu_models(prisma_client)
ptu_models: Final = (await _load_ptu_models(prisma_client)).models
days: Final = _backfill_window(ptu_models, end)
if not days:
@ -702,31 +724,68 @@ async def _deliver_alert(alert: "Callable[[str], Awaitable[None]] | None", messa
verbose_proxy_logger.error("PTU rollup: could not deliver the failed-charge alert: %s", exc)
def _prune_filter(*, date_str: str, cutoff: datetime, chunk: "tuple[str, ...] | None") -> "Mapping[str, object]":
"""One delete statement's predicate. An absent chunk leaves the sweep unbounded.
Returns a plain dict because the query builder serialises the mapping it is handed and
rejects a read-only view of one.
"""
return { # mutable-ok: prisma delete filter
"date": date_str,
"api_key": PTU_SENTINEL_API_KEY,
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
**({} if chunk is None else {"model": {"in": chunk}}), # mutable-ok: prisma membership filter
}
async def _prune_unrefreshed_sentinel_rows(
prisma_client: "PrismaClient",
*,
date_str: str,
run_started: datetime,
scanned_ids: frozenset[str] | None,
) -> None:
"""Delete the day's PTU sentinel rows this run did not refresh.
"""Delete the day's PTU sentinel rows this run looked at and did not refresh.
Every charge the run wrote bumps ``updated_at`` past ``run_started``, so anything
left below that mark is a (team, model) the current config no longer prices. The mark
is pulled back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come
from different hosts: a stale row is hours old, a concurrently written one is seconds
old, and the grace separates them without waiting on clocks agreeing. The
predicate reads only the row, never the caller's config snapshot, which is what
makes it safe to run twice, out of order, or beside another pod: a row written
after this run began is out of reach of its delete. Mirrors the retention predicate
``SpendLogCleanup`` deletes by."""
Two conditions, and a row survives unless it meets both. It must be stale: every
charge the run wrote bumps ``updated_at`` past ``run_started``, so anything left below
that mark is a (team, model) the current config no longer prices. The mark is pulled
back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come from
different hosts, and the grace separates a row that is hours old from one written
seconds ago without waiting on clocks agreeing.
A run that priced a deployment only its own host declares must also name the
deployments it scanned. Staleness alone is sufficient while every run derives its
charges from the same table, because then any two runs compute the same set, so a
database-only run still sweeps by timestamp exactly as it always has. Once one host's
charges come from a file the others cannot read, a row it never considered is not
evidence of anything, and deleting it drops a charge that host is responsible for.
Where the bound applies the ids go out in chunks, because each is one bind variable and
the server rejects a statement carrying more than 32767 of them, which a proxy holding
that many deployments would otherwise hit every night with no handler above here.
"""
cutoff: Final = run_started - timedelta(seconds=PTU_PRUNE_SKEW_GRACE_SECONDS)
await prisma_client.db.litellm_dailyteamspend.delete_many(
where={ # mutable-ok: prisma delete filter
"date": date_str,
"api_key": PTU_SENTINEL_API_KEY,
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
}
ordered: Final = () if scanned_ids is None else tuple(sorted(scanned_ids))
chunks: Final = (
(None,)
if scanned_ids is None
else tuple(
ordered[start : start + _PRUNE_ID_CHUNK_SIZE] for start in range(0, len(ordered), _PRUNE_ID_CHUNK_SIZE)
)
)
filters: Final = tuple(_prune_filter(date_str=date_str, cutoff=cutoff, chunk=chunk) for chunk in chunks)
deletions: Final = tuple(
[await prisma_client.db.litellm_dailyteamspend.delete_many(where=where) for where in filters]
)
deleted: Final = sum(deletions)
if deleted:
verbose_proxy_logger.info(
"PTU rollup for %s: pruned %s stale sentinel row(s) across %s deployment(s)",
date_str,
deleted,
"every" if scanned_ids is None else len(scanned_ids),
)
__all__ = (

View file

@ -216,6 +216,15 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d
return {}
def _sl_attribution_fallback(
standard_logging_payload: StandardLoggingPayload | None,
field: Literal["model_id", "model_group", "api_base", "custom_llm_provider"],
) -> str:
if standard_logging_payload is None:
return ""
return standard_logging_payload.get(field) or ""
def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload:
if kwargs is None:
kwargs = {}
@ -288,8 +297,15 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
): # use 'tags' from standard logging payload instead
request_tags = safe_dumps(standard_logging_payload["request_tags"])
_model_id: Final = metadata.get("model_info", {}).get("id", "")
_model_group: Final = metadata.get("model_group", "")
_model_id: Final = metadata.get("model_info", {}).get("id", "") or _sl_attribution_fallback(
standard_logging_payload, "model_id"
)
_model_group: Final = metadata.get("model_group", "") or _sl_attribution_fallback(
standard_logging_payload, "model_group"
)
_api_base: Final = litellm_params.get("api_base", "") or _sl_attribution_fallback(
standard_logging_payload, "api_base"
)
# Extract overhead from hidden_params if available
litellm_overhead_time_ms = None
@ -389,7 +405,11 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
# Extract agent_id for A2A requests (set directly on model_call_details)
agent_id: Final[str | None] = kwargs.get("agent_id") or metadata.get("agent_id")
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
custom_llm_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or None
)
raw_model: Final = cast(str, kwargs.get("model") or "")
model_name: Final = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
@ -414,13 +434,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
completion_tokens=usage.get("completion_tokens", standard_logging_completion_tokens),
request_tags=request_tags,
end_user=end_user_id or "",
api_base=litellm_params.get("api_base", ""),
api_base=_api_base,
model_group=_model_group,
model_id=_model_id,
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
agent_id=agent_id,
requester_ip_address=clean_metadata.get("requester_ip_address", None),
custom_llm_provider=kwargs.get("custom_llm_provider", ""),
custom_llm_provider=custom_llm_provider or "",
messages=_get_messages_for_spend_logs_payload(
standard_logging_payload=standard_logging_payload, metadata=metadata
),

View file

@ -516,6 +516,34 @@ def _failure_usage_to_lift(
return estimated_usage, 0.0
_EMPTY_LIFT: Final = MappingProxyType({})
def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]:
"""Failure-path callbacks run after ``litellm_logging_obj`` is popped from
request_data (it is not serialisable), so the caller merges these fields
onto request_data first: the first-handoff instant for preprocessing
latency, recovered or estimated usage for token counts, and the standard
logging object for deployment attribution on failed-request spend logs."""
_logging_obj: Final = request_data.get("litellm_logging_obj")
if _logging_obj is None:
return _EMPTY_LIFT
_model_call_details: Final = getattr(_logging_obj, "model_call_details", {})
_first_handoff: Final = _model_call_details.get("first_api_call_start_time")
_usage_to_lift: Final = _failure_usage_to_lift(
model_call_details=_model_call_details,
request_body=request_data,
dispatched=_first_handoff is not None,
)
_entries: Final = (
("first_api_call_start_time", _first_handoff),
("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]),
("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)),
("standard_logging_object", _model_call_details.get("standard_logging_object")),
)
return MappingProxyType({key: value for key, value in _entries if value is not None})
@dataclass(frozen=True)
class _CallbackCapabilities:
"""Cached per-hook capability flags derived from ``litellm.callbacks``.
@ -2280,6 +2308,11 @@ class ProxyLogging:
)
)
# Auth and pass-through failure bodies are unstripped client input, and
# the logging handler below flattens body keys into model_call_details,
# so drop the key before it can masquerade as the built payload.
request_data.pop("standard_logging_object", None)
### LOGGING ###
if self._is_proxy_only_llm_api_error(
original_exception=original_exception,
@ -2293,29 +2326,7 @@ class ProxyLogging:
original_exception=original_exception,
)
# Lift the first-handoff instant onto request_data (top-level
# internal key, not metadata) so failure-path callbacks can still
# compute preprocessing latency after the logging object is popped.
_logging_obj: Final = request_data.get("litellm_logging_obj")
if _logging_obj is not None:
_model_call_details: Final = getattr(_logging_obj, "model_call_details", {})
_first_handoff: Final = _model_call_details.get("first_api_call_start_time")
if _first_handoff is not None:
request_data["first_api_call_start_time"] = _first_handoff
# Lift recovered partial-stream usage, or an estimated input-side
# usage for a dispatched failure, onto request_data so the
# failure-path spend callbacks (which run after the logging object
# is popped) record real token counts instead of zero.
_usage_to_lift: Final = _failure_usage_to_lift(
model_call_details=_model_call_details,
request_body=request_data,
dispatched=_first_handoff is not None,
)
if _usage_to_lift is not None:
_lifted_usage, _lifted_cost = _usage_to_lift
request_data["combined_usage_object"] = _lifted_usage
request_data["response_cost"] = _lifted_cost
request_data.update(_failure_fields_to_lift(request_data))
# Remove before callbacks iterate — not serialisable
request_data.pop("litellm_logging_obj", None)

View file

@ -64,6 +64,7 @@ from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.ptu_pricing import zeroed_ptu_pricing
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
)
@ -7695,7 +7696,16 @@ class Router:
- None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params)
"""
try:
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**_litellm_params)
zeroed_pricing: Final = (
zeroed_ptu_pricing(_model_info, _litellm_params) if _model_info.get("db_model") is not True else None
)
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
**(
_litellm_params
if zeroed_pricing is None
else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
)
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
deployment = Deployment(
**deployment_info,
@ -11191,6 +11201,8 @@ class Router:
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
if pre_routing_hook_response.litellm_params:
request_kwargs.update(pre_routing_hook_response.litellm_params)
#########################################################
# Resolve the strategy and logger AFTER the pre-routing hook, since
@ -11300,6 +11312,8 @@ class Router:
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
if pre_routing_hook_response.litellm_params:
request_kwargs.update(pre_routing_hook_response.litellm_params)
# 2. Get healthy deployments
healthy_deployments: Final = await self.async_get_healthy_deployments(

View file

@ -53,6 +53,21 @@ model_list:
REASONING: o1-preview
```
Each tier can also use a model entry with request parameter overrides. A tier value may be
a model string, a single object, or a list mixing strings and objects. Object entries must
contain a model name and may contain any LiteLLM request parameters. The model name must
still resolve to a deployment in `model_list`; this configuration does not create one
```yaml
tiers:
COMPLEX: opus
REASONING:
- model_name: opus
litellm_params:
reasoning_effort: xhigh
- abc
```
### Renaming the tiers
`tier_labels` puts your own vocabulary on the four tiers:

View file

@ -30,6 +30,7 @@ from litellm.constants import EMPTY_MAPPING, RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
@ -663,6 +664,35 @@ class ClassificationOutcome(NamedTuple):
classifier_cost: float | None = None
class _SessionAffinityPin(NamedTuple):
model: str
tier: ComplexityTier | None
def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None:
if isinstance(value, str):
return _SessionAffinityPin(model=value, tier=None)
parts: Final[tuple[object, object] | None] = (
(value.get("model"), value.get("tier"))
if isinstance(value, Mapping)
else (value[0], value[1])
if isinstance(value, (list, tuple)) and len(value) == 2
else None
)
if parts is None:
return None
model, tier_value = parts
if not isinstance(model, str):
return None
tier: Final = ComplexityTier(tier_value) if isinstance(tier_value, str) else None
return _SessionAffinityPin(model=model, tier=tier)
def _session_affinity_cache_value(model: str, tier: ComplexityTier | str | None) -> Mapping[str, str | None]:
tier_value: Final = _tier_name(tier) if tier is not None else None
return {"model": model, "tier": tier_value} # mutable-ok: cache requires JSON mapping
class ComplexityRouter(CustomLogger):
"""
Complexity router that classifies requests and routes to appropriate models.
@ -1078,6 +1108,7 @@ class ComplexityRouter(CustomLogger):
classifier_model: str | None = None,
classifier_cost: float | None = None,
conversation_continuing: bool = True,
tier_litellm_params: Mapping[str, object] | None = None,
) -> StandardLoggingRoutingDecision:
"""Assemble the per-request provenance record for this router's decision.
@ -1127,6 +1158,10 @@ class ComplexityRouter(CustomLogger):
decision["classifier_model"] = classifier_model
if classifier_cost is not None:
decision["classifier_cost"] = classifier_cost
if tier_litellm_params:
masked_tier_litellm_params: Final = mask_credentials_in_payload(tier_litellm_params)
if isinstance(masked_tier_litellm_params, Mapping):
decision["tier_litellm_params"] = masked_tier_litellm_params
return decision
async def aclassify(
@ -1457,6 +1492,13 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"No model configured for tier {tier_key} and no default_model set")
def _litellm_params_for_model(self, tier: ComplexityTier | str | None, model: str) -> Mapping[str, object]:
if tier is None:
return MappingProxyType({})
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ())
entry: Final = next((candidate for candidate in entries if candidate.model_name == model), None)
return entry.litellm_params if entry is not None else MappingProxyType({})
@staticmethod
def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str:
if isinstance(model, str):
@ -2068,9 +2110,10 @@ class ComplexityRouter(CustomLogger):
cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None
if cache_key is not None:
pinned_model: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
if isinstance(pinned_model, str):
routed_model: str | None = pinned_model
pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
pinned_pin: Final = _parse_session_affinity_pin(pinned_value)
if pinned_pin is not None:
routed_model: str | None = pinned_pin.model
pin_escalation_keyword: str | None = None
if self.escalation_keywords:
user_message: Final = (
@ -2079,16 +2122,21 @@ class ComplexityRouter(CustomLogger):
if user_message is not None:
pin_escalation_keyword = self._matched_escalation_keyword(user_message)
if pin_escalation_keyword is not None:
routed_model = self._escalated_pin(pinned_model)
routed_model = self._escalated_pin(pinned_pin.model)
if routed_model is not None:
escalated: Final = routed_model != pinned_model
escalated: Final = routed_model != pinned_pin.model
resolved_pin_tier: Final = (
pinned_pin.tier
if not escalated and pinned_pin.tier is not None
else self._tier_for_model(routed_model)
)
# The floor outranks the pin because plan mode is a transient state of the
# session, not a request to move it: the turns carrying the sentinel route at
# the floor, and the stored pin deliberately keeps the session's own model so
# the first turn after plan mode exits auto-routes exactly as it would have.
# Escalation is the opposite on purpose -- an explicit ask to re-pin higher.
pin_plan_sentinel: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages)
pinned_tier: Final = self._tier_for_model(routed_model) if pin_plan_sentinel is not None else None
pinned_tier: Final = resolved_pin_tier if pin_plan_sentinel is not None else None
plan_floored: Final = (
pinned_tier is not None and self._apply_plan_mode_floor(pinned_tier) != pinned_tier
)
@ -2099,7 +2147,7 @@ class ComplexityRouter(CustomLogger):
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=session_model,
value=_session_affinity_cache_value(session_model, resolved_pin_tier),
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
@ -2118,19 +2166,23 @@ class ComplexityRouter(CustomLogger):
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, routed_model=%s", cause, routed_model
)
routed_pin_tier: Final = self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
has_original_messages: Final = messages is not None and len(messages) > 0
return self._with_session_deployment_affinity(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=self._tier_for_model(routed_model),
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
),
)
)
@ -2157,7 +2209,10 @@ class ComplexityRouter(CustomLogger):
if pinnable and cache_key is not None and response is not None:
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=response.model,
value=_session_affinity_cache_value(
response.model,
response.routing_decision.get("tier") if response.routing_decision is not None else None,
),
ttl=self.config.session_affinity_ttl_seconds,
)
return self._with_session_deployment_affinity(response)
@ -2271,6 +2326,7 @@ class ComplexityRouter(CustomLogger):
)
keyword_plan_floored: Final = routed_tier != escalated_tier
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
keyword_tier_litellm_params: Final = self._litellm_params_for_model(routed_tier, routed_model)
keyword_cause: Final[RoutingDecisionCause] = (
"plan_mode"
if keyword_plan_floored
@ -2286,6 +2342,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=keyword_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,
@ -2294,6 +2351,7 @@ class ComplexityRouter(CustomLogger):
matched_keyword=plan_mode_sentinel if keyword_plan_floored else override.matched_keyword,
escalation_keyword=escalation_keyword,
escalated=keyword_escalated,
tier_litellm_params=keyword_tier_litellm_params,
),
)
@ -2380,6 +2438,7 @@ class ComplexityRouter(CustomLogger):
routed_model,
)
tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model)
classifier_model: Final = (
self.config.classifier_llm_config.model
if outcome.cause == "llm_classifier" and self.config.classifier_llm_config is not None
@ -2405,6 +2464,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,
@ -2417,5 +2477,6 @@ class ComplexityRouter(CustomLogger):
escalated=escalated,
classifier_model=classifier_model,
classifier_cost=outcome.classifier_cost,
tier_litellm_params=tier_litellm_params,
),
)

View file

@ -5,10 +5,12 @@ Contains default keyword lists, weights, tier boundaries, and configuration clas
All values are configurable via proxy config.yaml.
"""
from collections.abc import Mapping
from enum import Enum
from typing import Final, Literal
from types import MappingProxyType
from typing import Annotated, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_serializer, field_validator, model_validator
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
@ -159,6 +161,44 @@ class ReminderMarkerPair(BaseModel):
return self
class ComplexityTierModel(BaseModel):
model_config = ConfigDict(frozen=True)
model_name: str
litellm_params: Annotated[Mapping[str, object], SkipValidation()] = Field(
default_factory=lambda: MappingProxyType({})
)
@field_validator("litellm_params", mode="before")
@classmethod
def _freeze_litellm_params(cls, value: Mapping[str, object]) -> Mapping[str, object]:
return MappingProxyType(dict(value))
@field_serializer("litellm_params")
def _serialize_litellm_params(self, value: Mapping[str, object]) -> Mapping[str, object]:
return dict(value) # mutable-ok: Pydantic JSON serialization requires a concrete mapping
def _normalize_tier_entries(
raw_value: object,
tier: str,
) -> tuple[str | list[str], tuple[ComplexityTierModel, ...]]:
raw_entries: Final = raw_value if isinstance(raw_value, (list, tuple)) else (raw_value,)
entries: Final = tuple(
ComplexityTierModel(model_name=entry) if isinstance(entry, str) else ComplexityTierModel.model_validate(entry)
for entry in raw_entries
)
model_names: Final = tuple(entry.model_name for entry in entries)
if len(model_names) != len(frozenset(model_names)):
raise ValueError(f"tier {tier} contains duplicate model_name values; each pool entry needs distinct parameters")
normalized: Final = (
entries[0].model_name
if not isinstance(raw_value, (list, tuple))
else list(model_names) # mutable-ok: config.tiers must preserve its existing list contract
)
return normalized, entries
# ─── Default Keyword Lists ───
# Note: Keywords should be full words/phrases to avoid substring false positives.
# The matching logic uses word boundary detection for single-word keywords.
@ -425,6 +465,9 @@ class ComplexityRouterConfig(BaseModel):
"A list is randomly picked from when adaptive=False, and used as a soft-floor home pool when adaptive=True"
),
)
tier_model_configs: Mapping[str, tuple[ComplexityTierModel, ...]] = Field(
default_factory=dict,
)
tier_definitions: tuple[TierDefinition, ...] | None = Field(
default=None,
@ -777,6 +820,55 @@ class ComplexityRouterConfig(BaseModel):
coerced[key] = item
return coerced
@model_validator(mode="before")
@classmethod
def _normalize_tier_model_configs(cls, value: object) -> object:
if not isinstance(value, dict):
return value
raw_tiers: Final = value.get("tiers")
if not isinstance(raw_tiers, dict):
return value
existing_configs: Final = value.get("tier_model_configs")
normalized_entries: Final = MappingProxyType(
{tier: _normalize_tier_entries(raw_value, tier) for tier, raw_value in raw_tiers.items()}
)
normalized_tiers: Final = MappingProxyType(
{tier: normalized for tier, (normalized, _) in normalized_entries.items()}
)
incoming_params: Final = (
MappingProxyType(
{
(tier, entry.model_name): entry.litellm_params
for tier, entries in existing_configs.items()
for entry in (ComplexityTierModel.model_validate(item) for item in entries)
}
)
if isinstance(existing_configs, dict)
else MappingProxyType({})
)
tier_model_configs: Final = MappingProxyType(
{
tier: tuple(
entry.model_copy(
update=MappingProxyType(
{
"litellm_params": incoming_params.get((tier, entry.model_name), entry.litellm_params),
}
)
)
for entry in entries
)
for tier, (_, entries) in normalized_entries.items()
if any(entry.litellm_params for entry in entries)
or (isinstance(existing_configs, dict) and tier in existing_configs)
}
)
return { # mutable-ok: Pydantic before-validator requires a concrete mapping
**value,
"tiers": normalized_tiers,
"tier_model_configs": tier_model_configs,
}
@field_validator("escalation_keywords")
@classmethod
def _normalize_escalation_keywords(cls, value: list[str] | None) -> list[str] | None:

View file

@ -4,6 +4,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
import datetime
import enum
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
@ -897,6 +898,7 @@ class PreRoutingHookResponse(BaseModel):
messages: list[dict[str, Any]] | None
routing_decision: StandardLoggingRoutingDecision | None = None
session_affinity_ttl_seconds: int | None = None
litellm_params: Mapping[str, object] | None = None
_PreRoutingStrategyT_co = TypeVar("_PreRoutingStrategyT_co", covariant=True)

View file

@ -2843,6 +2843,7 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
conversation_continuing: bool
savings_baseline_model: str
savings_baseline_deployment_id: str
tier_litellm_params: Mapping[str, object] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
# Fields whose values quote the caller's prompt. Dropped when an operator turns message
@ -2868,6 +2869,7 @@ DERIVED_ROUTING_DECISION_FIELDS: Final[frozenset[str]] = frozenset(
"conversation_continuing",
"savings_baseline_model",
"savings_baseline_deployment_id",
"tier_litellm_params",
}
)

View file

@ -17,6 +17,7 @@ from __future__ import annotations
import json
import os
import re
import time
from datetime import datetime, timedelta, timezone
from typing import Callable
@ -57,7 +58,7 @@ from e2e_http import (
unwrap,
)
from lifecycle import ResourceManager
from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogRow
from models import KeyGenerateBody, KeyMetadata, LiteLLMParamsBody, SpendLogRow
pytestmark = pytest.mark.e2e
@ -685,6 +686,149 @@ class TestBatchRateLimitErrorMapping:
)
BATCH_ENQUEUED_HEADROOM_TOKENS = 100_000
_BATCH_REQUIRES_TOKENS = re.compile(r"Batch requires (\d+) tokens")
class TestBatchEnqueuedTokenLimit:
"""Opt-in enqueued-token allowance governs batch submission instead of RPM/TPM.
A key whose metadata carries batch_enqueued_token_limit reserves the batch's
token estimate against that allowance at create time: per-minute limits no
longer gate batch submission, exhausting the allowance rejects the create
before it reaches the provider, and cancelling a running batch refunds its
reservation so blocked submissions go through again (LIT-5273).
"""
def _upload_batch_file(
self, client: BatchClient, resources: ResourceManager, key: str
) -> FileObject:
file = unwrap(
client.upload_file(
content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES),
form=FileUploadForm(purpose="batch"),
model=OPENAI_BATCH_MODEL,
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
return file
def _generate_enqueued_key(
self,
client: BatchClient,
resources: ResourceManager,
*,
limit: int,
marker: str,
rpm_limit: int | None = None,
) -> str:
key = client.proxy.generate_key(
KeyGenerateBody(
models=[],
rpm_limit=rpm_limit,
user_id=f"e2e-batch-enq-{marker}-{unique_marker()}",
metadata=KeyMetadata(batch_enqueued_token_limit=limit),
)
)
resources.defer(lambda: client.proxy.delete_key(key))
return key
@pytest.mark.covers(
"quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm",
exercised_on=["batches"],
)
def test_enqueued_allowance_accepts_batch_over_key_rpm(
self, client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:
key = self._generate_enqueued_key(
client,
resources,
limit=BATCH_ENQUEUED_HEADROOM_TOKENS,
marker="rpm",
rpm_limit=BATCH_RL_RPM_LIMIT,
)
file = self._upload_batch_file(client, resources, key)
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert created.status_code != 429, (
f"enqueued-token allowance must govern batch submission instead of the "
f"key RPM ({BATCH_RL_RPM_LIMIT} < {BATCH_RL_REQUEST_LINES} rows); "
f"got 429: {created.body[:400]}"
)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
@pytest.mark.covers(
"quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted",
exercised_on=["batches"],
)
@pytest.mark.covers(
"quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel",
exercised_on=["batches"],
)
def test_exhausted_allowance_blocks_until_cancel_refunds(
self, client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:
sizing_key = self._generate_enqueued_key(
client, resources, limit=1, marker="size"
)
sizing_file = self._upload_batch_file(client, resources, sizing_key)
sized = client.create_batch(
body=BatchCreateBody(input_file_id=sizing_file.id), key=sizing_key
)
assert sized.status_code == 429, (
f"a 1-token allowance must reject any batch before it reaches the "
f"provider, got {sized.status_code}: {sized.body[:400]}"
)
assert "batch enqueued token limit exceeded" in sized.body.lower(), (
f"429 body must name the enqueued token limit, got: {sized.body[:400]}"
)
requires = _BATCH_REQUIRES_TOKENS.search(sized.body)
assert requires is not None, (
f"429 body must report the batch token requirement so callers can size "
f"allowances, got: {sized.body[:400]}"
)
batch_tokens = int(requires.group(1))
assert batch_tokens > 1
key = self._generate_enqueued_key(
client, resources, limit=batch_tokens + batch_tokens // 2, marker="refund"
)
file = self._upload_batch_file(client, resources, key)
first = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(first)
first_batch = BatchObject.model_validate_json(first.body)
resources.defer(quietly(lambda: client.cancel_batch(first_batch.id, key=key)))
blocked = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert blocked.status_code == 429, (
f"second batch must not fit the remaining allowance while the first is "
f"enqueued, got {blocked.status_code}: {blocked.body[:400]}"
)
assert "batch enqueued token limit exceeded" in blocked.body.lower(), (
f"429 body must name the enqueued token limit, got: {blocked.body[:400]}"
)
cancelled = cancel_batch(client, first_batch.id, key=key, provider=None)
assert cancelled.status in {"cancelling", "cancelled"}, (
f"cancel must reach a cancel state for the refund to fire, "
f"got {cancelled.status}"
)
retried = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert retried.status_code != 429, (
f"cancelling the first batch must refund its reservation so the retry "
f"fits the allowance, got 429: {retried.body[:400]}"
)
require_successful_call(retried)
retry_batch = BatchObject.model_validate_json(retried.body)
resources.defer(quietly(lambda: client.cancel_batch(retry_batch.id, key=key)))
ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"

View file

@ -2,6 +2,9 @@
# litellm/proxy/hooks/ + litellm/proxy/auth/auth_checks.py + litellm/proxy/spend_tracking/.
- {id: quota_management.ratelimit.rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: rpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces RPM per key/team/model; 429 on breach"}
- {id: quota_management.ratelimit.batch_rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_rpm, assertions: [blocks_over_limit], exercised_on: [batches], source: "batch_rate_limiter.py", rationale: "Batch create that exceeds key RPM returns mapped 429 with retry-after"}
- {id: quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [accepts_over_rpm], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Key with an enqueued-token allowance submits a batch whose row count exceeds its RPM and the create is accepted"}
- {id: quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [blocks_when_exhausted], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Batch create is rejected with a 429 naming the enqueued token limit once the allowance cannot fit the file"}
- {id: quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [refunds_on_cancel], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Cancelling a running batch returns its reserved tokens so a previously blocked submission succeeds"}
- {id: quota_management.ratelimit.tpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces TPM per key/team/model; 429 on breach"}
- {id: quota_management.ratelimit.tpm.excludes_cached_tokens, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [excludes_cached_tokens], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py:_get_total_tokens_from_usage", rationale: "Cached prompt tokens must not count toward TPM (LIT-1930)"}
- {id: quota_management.ratelimit.redis_backed.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: redis_backed, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py", rationale: "With Redis configured, RPM still enforces 429 across the shared limiter path customers run multi-replica"}

View file

@ -46,6 +46,7 @@ class KeyLoggingCallback(BaseModel):
class KeyMetadata(BaseModel):
logging: list[KeyLoggingCallback] | None = None
priority: str | None = None
batch_enqueued_token_limit: int | None = None
class ObjectPermission(BaseModel):

View file

@ -12,23 +12,21 @@ export async function createMcpServer(page: PwPage, url: string): Promise<string
await expect(discovery).toBeVisible({ timeout: 5_000 });
await discovery.getByRole("button", { name: /Custom Server/i }).click();
const formModal = page.locator(".ant-modal:visible").filter({ hasText: "MCP Server Name" });
const formModal = page.getByRole("dialog").filter({ hasText: "MCP Server Name" });
await expect(formModal).toBeVisible({ timeout: 5_000 });
// validateMCPServerName rejects spaces and hyphens; the worker index avoids a same-millisecond collision.
const name = `e2e_mcp_${process.env.TEST_WORKER_INDEX ?? "0"}_${Date.now()}`;
await formModal.locator('input[id="server_name"]').fill(name);
await formModal.getByLabel("MCP Server Name").fill(name);
const transportField = formModal.locator(".ant-form-item", { hasText: "Transport Type" });
await transportField.locator(".ant-select").click();
await page.locator(".ant-select-dropdown:visible").getByText("Streamable HTTP").click();
// Select popups are portaled to the body, so the option lookup is page-scoped, not modal-scoped.
await formModal.getByRole("combobox", { name: "Transport Type" }).click();
await page.getByRole("option", { name: "Streamable HTTP" }).click();
await formModal.locator('input[id="url"]').fill(url);
await formModal.getByLabel("MCP Server URL").fill(url);
// The auth_type Form.Item has no label prop, so anchor on the enclosing Collapse panel.
const authSection = formModal.locator(".ant-collapse-item", { hasText: /^Authentication/ });
await authSection.locator(".ant-form-item").first().locator(".ant-select").click();
await page.locator(".ant-select-dropdown:visible").getByText("None", { exact: true }).click();
await formModal.getByRole("combobox", { name: "Authentication", exact: true }).click();
await page.getByRole("option", { name: "None", exact: true }).click();
await formModal.getByRole("button", { name: /^Add MCP Server$/ }).click();
await expect(page.getByText("MCP Server created successfully").first()).toBeVisible({ timeout: 15_000 });

View file

@ -30,31 +30,26 @@ test.describe("MCP Servers", () => {
await expect(discovery).toBeVisible({ timeout: 5_000 });
await discovery.getByRole("button", { name: /Custom Server/i }).click();
const formModal = page.locator(".ant-modal:visible").filter({ hasText: "MCP Server Name" });
const formModal = page.getByRole("dialog").filter({ hasText: "MCP Server Name" });
await expect(formModal).toBeVisible({ timeout: 5_000 });
// Name — no spaces or hyphens per validateMCPServerName
const uniqueName = `e2e_mcp_${Date.now()}`;
createdServerName = uniqueName;
await formModal.locator('input[id="server_name"]').fill(uniqueName);
await formModal.getByLabel("MCP Server Name").fill(uniqueName);
// Transport: Streamable HTTP — the only value the proxy actually accepts is "http"
const transportField = formModal.locator(".ant-form-item", { hasText: "Transport Type" });
await transportField.locator(".ant-select").click();
await page.locator(".ant-select-dropdown:visible").getByText("Streamable HTTP").click();
// Transport: Streamable HTTP — the only value the proxy actually accepts is "http".
// Select popups are portaled to the body, so the option lookup is page-scoped.
await formModal.getByRole("combobox", { name: "Transport Type" }).click();
await page.getByRole("option", { name: "Streamable HTTP" }).click();
// URL — use a fake URL; the form just persists it, it doesn't have to be reachable
await formModal.locator('input[id="url"]').fill("https://e2e-fake-mcp.test.local/mcp");
await formModal.getByLabel("MCP Server URL").fill("https://e2e-fake-mcp.test.local/mcp");
// Authentication: None
// The auth_type Form.Item has no label prop (CreateMCPServer.tsx), so
// it can't be anchored by label text. Scope via the enclosing Collapse
// panel ("Authentication") instead — that anchor is stable even if the
// placeholder copy changes.
const authSection = formModal.locator(".ant-collapse-item", { hasText: /^Authentication/ });
const authField = authSection.locator(".ant-form-item").first();
await authField.locator(".ant-select").click();
await page.locator(".ant-select-dropdown:visible").getByText("None", { exact: true }).click();
// Authentication: None. "Authentication" is exact so it can't also match the
// "Authentication Value" field that some auth types reveal below it.
await formModal.getByRole("combobox", { name: "Authentication", exact: true }).click();
await page.getByRole("option", { name: "None", exact: true }).click();
// Submit
await formModal.getByRole("button", { name: /^Add MCP Server$/ }).click();

View file

@ -21,15 +21,22 @@ async function findDeploymentByName(page: PlaywrightPage, modelName: string): Pr
return body.data.find((row) => row.model_name === modelName);
}
/** Anchors a substring match to the whole string, escaping regex metacharacters. */
const exactly = (text: string): RegExp => new RegExp(`^${text.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}$`);
/**
* Helper to select a provider from the Add Model form dropdown.
* Helper to select a provider from the Add Model form dropdown. The field is a
* searchable combobox: it only opens on click, typing filters the list, and the
* option has to be picked explicitly because nothing is highlighted by default.
* Options are matched on their visible text, not their accessible name, which
* also carries the provider logo's alt text ("Anthropic logo Anthropic").
*/
async function selectProvider(page: any, providerName: string) {
const providerDropdown = page.getByRole("combobox", { name: /Provider/i });
async function selectProvider(page: PlaywrightPage, providerName: string) {
const providerDropdown = page.getByRole("combobox", { name: "Provider", exact: true });
await providerDropdown.click();
await providerDropdown.fill(providerName);
await page.waitForTimeout(1000);
await providerDropdown.press("Enter");
await page.waitForTimeout(2000);
await page.getByRole("option").filter({ hasText: exactly(providerName) }).click();
await expect(providerDropdown).toHaveValue(providerName);
}
test.describe("Add Model", () => {
@ -64,11 +71,10 @@ test.describe("Add Model", () => {
await selectProvider(page, "Anthropic");
// The model field should be a multi-select dropdown; click to open it
const modelDropdown = page.locator(".ant-select-selection-overflow").first();
await modelDropdown.click();
await page.getByRole("combobox", { name: "Select models" }).click();
// Verify provider-specific models are listed
await expect(page.getByTitle("claude-haiku-4-5", { exact: true })).toBeVisible();
await expect(page.getByRole("option", { name: "claude-haiku-4-5", exact: true })).toBeVisible();
});
test("Edit team model TPM and RPM limits", async ({ page }) => {
@ -156,14 +162,14 @@ test.describe("Add Model", () => {
await page.getByRole("tab", { name: "Add Model" }).click();
// Labels come from /public/providers/fields, not the frontend Providers enum, and the two differ.
await selectProvider(page, "OpenAI-Compatible Endpoints");
await selectProvider(page, "OpenAI-Compatible Endpoints (Together AI, etc.)");
const publicName = `e2e-ui-added-${Date.now()}`;
uiAddedModelName = publicName;
// The model picker's "custom" entry reveals the free-text name field.
await page.locator(".ant-select-selection-overflow").first().click();
await page.locator(".ant-select-dropdown:visible").getByText("Custom Model Name (Enter below)").click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "Custom Model Name (Enter below)" }).click();
await page.keyboard.press("Escape");
await page.getByPlaceholder("Enter custom model name").fill(publicName);
@ -177,8 +183,8 @@ test.describe("Add Model", () => {
await expect(page.getByTestId("connection-success-msg")).toBeVisible({ timeout: 30_000 });
// The modal swallows the Add click. Scope to the footer: the dismiss X is also named "Close".
const resultsModal = page.locator(".ant-modal:visible").filter({ hasText: "Connection Test Results" });
await resultsModal.locator(".ant-modal-footer").getByRole("button", { name: "Close" }).click();
const resultsModal = page.getByRole("dialog", { name: "Connection Test Results" });
await resultsModal.locator('[data-slot="dialog-footer"]').getByRole("button", { name: "Close" }).click();
await expect(resultsModal).toBeHidden({ timeout: 5_000 });
const created = await captureRequestBody(page, { method: "POST", urlIncludes: "/model/new" }, async () => {
@ -213,9 +219,8 @@ test.describe("Add Model", () => {
await selectProvider(page, "Anthropic");
// Select model: claude-haiku-4-5
const modelDropdown = page.locator(".ant-select-selection-overflow").first();
await modelDropdown.click();
await page.getByTitle("claude-haiku-4-5", { exact: true }).click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "claude-haiku-4-5", exact: true }).click();
await page.keyboard.press("Escape");
// Enter bad API key
@ -239,9 +244,8 @@ test.describe("Add Model", () => {
await selectProvider(page, "Anthropic");
// Select model: claude-haiku-4-5
const modelDropdown = page.locator(".ant-select-selection-overflow").first();
await modelDropdown.click();
await page.getByTitle("claude-haiku-4-5", { exact: true }).click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "claude-haiku-4-5", exact: true }).click();
await page.keyboard.press("Escape");
// Enter any API key
@ -315,18 +319,15 @@ test.describe("Add Model", () => {
await selectProvider(page, "Cohere");
const modelDropdown = page.locator(".ant-select-selection-overflow").first();
await modelDropdown.click();
const wildcardOption = page.getByTitle(/All .* Models \(Wildcard\)/);
await wildcardOption.click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: /All .* Models \(Wildcard\)/ }).click();
await page.keyboard.press("Escape");
const apiKeyInput = page.locator('input[type="password"]').first();
await apiKeyInput.fill("sk-any-key-for-team-byok-test");
// Flip the Team-BYOK switch on (Form.Item label "Team-BYOK Model")
const teamByokRow = page.locator(".ant-form-item", { hasText: "Team-BYOK Model" });
await teamByokRow.getByRole("switch").click();
// Flip the Team-BYOK switch on; the Switch carries its own aria-label.
await page.getByRole("switch", { name: "Team-BYOK Model" }).click();
// TeamDropdown options show the alias above the team id, so match on the id line by text.
const teamDropdown = page.getByTestId("team-dropdown").getByRole("combobox");
@ -376,10 +377,8 @@ test.describe("Add Model", () => {
await selectProvider(page, "Cohere");
// Select All Cohere Models (Wildcard)
const modelDropdown = page.locator(".ant-select-selection-overflow").first();
await modelDropdown.click();
const wildcardOption = page.getByTitle(/All .* Models \(Wildcard\)/);
await wildcardOption.click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: /All .* Models \(Wildcard\)/ }).click();
await page.keyboard.press("Escape");
// Enter any API key

View file

@ -41,7 +41,7 @@ test.describe("Edit LLM credential", () => {
await row.getByTestId(`credential-actions-${credentialName}`).click();
await page.getByTestId("credential-action-edit").click();
const modal = page.locator(".ant-modal-content").filter({ hasText: "Edit Credential" });
const modal = page.getByRole("dialog", { name: "Edit Credential" });
await expect(modal).toBeVisible({ timeout: 10_000 });
const apiKeyField = modal.locator("#api_key");

View file

@ -45,9 +45,9 @@ test.describe("Proxy Admin - Keys", () => {
await page.keyboard.type(E2E_TEAM_CRUD_ALIAS);
await page.locator('[data-slot="combobox-content"]:visible').getByText(E2E_TEAM_CRUD_ALIAS).first().click();
// Select models
await page.locator(".ant-select-selection-overflow").click();
await page.locator(".ant-select-dropdown:visible").getByText("All Team Models").click();
// Select models — the popup is portaled to the body, so scope options to the page.
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "All Team Models", exact: true }).click();
await page.keyboard.press("Escape");
// Submit
@ -86,7 +86,7 @@ test.describe("Proxy Admin - Keys", () => {
// Scope to the modal — the Regenerate button has an icon whose aria-label
// ("sync") is concatenated into the button's accessible name, and the
// "Regenerate Key" button is still in the DOM behind the modal.
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Regenerate Virtual Key" });
await modal.getByRole("button", { name: /Regenerate/ }).click();
// Success view shows a Copy button in the footer (text varies between modal versions)
@ -198,8 +198,8 @@ test.describe("Proxy Admin - Keys", () => {
// Select models — open the multi-select and pick the all-models meta-option.
// With no team selected the modal offers "All Proxy Models"; the team-scoped
// "All Team Models" option only appears once a team is picked.
await page.locator(".ant-select-selection-overflow").click();
await page.locator(".ant-select-dropdown:visible").getByText("All Proxy Models").click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "All Proxy Models", exact: true }).click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
@ -221,17 +221,10 @@ test.describe("Proxy Admin - Keys", () => {
const keyName = `e2e-admin-specific-${Date.now()}`;
await page.getByLabel(/Key Name/).fill(keyName);
// Open the model multi-select and pick a single specific model. Use
// getByRole("option", ...) to avoid the strict-mode collision between
// the option container and its inner text node.
// Open the model multi-select and pick a single specific model.
const modelName = "fake-openai-gpt-4";
await page.locator(".ant-select-selection-overflow").click();
const option = page.locator(".ant-select-dropdown:visible").getByRole("option", { name: modelName, exact: true });
await option.waitFor({ state: "attached" });
// Dispatch the click via the DOM — antd's dropdown can render the option
// off-viewport during the open animation, which trips Playwright's
// visibility/stability checks. The click handler fires regardless.
await option.evaluate((el: HTMLElement) => el.click());
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: modelName, exact: true }).click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
@ -242,7 +235,7 @@ test.describe("Proxy Admin - Keys", () => {
// verify it can call /chat/completions for the model it was scoped to.
// The mock LLM server (fixtures/mock_llm_server/server.py) replies with
// a fixed "This is a mock response." body.
const apiKey = (await page.locator(".ant-modal:visible pre").innerText()).trim();
const apiKey = (await page.getByRole("dialog", { name: "Save your Key" }).locator("pre").innerText()).trim();
expect(apiKey).toMatch(/^sk-/);
const response = await page.request.post("/chat/completions", {

View file

@ -41,11 +41,12 @@ test.describe("Proxy Admin - Teams", () => {
.click();
// Wait for the Create Team modal
const dialog = page.locator(".ant-modal:visible");
const dialog = page.getByRole("dialog", { name: "Create Team" });
await expect(dialog).toBeVisible({ timeout: 5_000 });
// Fill Team Name — the input has id="team_alias"
await dialog.locator("#team_alias").fill(uniqueAlias);
// Fill Team Name — FormField derives the control id from React.useId(), so
// the input is only addressable by its label or its test id.
await dialog.getByTestId("team-name-input").fill(uniqueAlias);
// Select models — the models multi-select is inside the modal. Its popup is
// portaled to the body, so scope the option lookup to the page, not the dialog.
@ -75,7 +76,7 @@ test.describe("Proxy Admin - Teams", () => {
await page.getByRole("button", { name: /Add Member/i }).click();
// Wait for Add Team Member modal
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Add Team Member" });
await expect(modal).toBeVisible({ timeout: 5_000 });
// The email field is a Select — type to search, then select from dropdown
@ -112,7 +113,7 @@ test.describe("Proxy Admin - Teams", () => {
await page.getByTestId("edit-member").first().click();
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Edit Member" });
await expect(modal).toBeVisible({ timeout: 5_000 });
await modal.getByRole("button", { name: /Save Changes/i }).click();
@ -155,7 +156,7 @@ test.describe("Proxy Admin - Teams", () => {
await page.getByTestId("edit-member").first().click();
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Edit Member" });
await expect(modal).toBeVisible({ timeout: 5_000 });
await modal.getByRole("button", { name: /Save Changes/i }).click();

View file

@ -67,29 +67,25 @@ test.describe("Router Settings - Fallbacks", () => {
await page.getByRole("button", { name: /Add Fallbacks/i }).click();
await modelsLoaded;
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Configure Model Fallbacks" });
await expect(modal).toBeVisible({ timeout: 5_000 });
// FallbackGroupConfig.tsx renders both selects with `showSearch`. The
// most stable interaction is: click to open + focus, type the model name to
// narrow the listbox to a single highlighted option, then press Enter.
// Verify each selection landed by watching the dialog's own state transition
// (the tab title updates to the picked primary; the fallback chain list
// populates) rather than by asserting on the dropdown popup, which sits in
// a custom getPopupContainer and is awkward to scope reliably.
const primarySelect = modal.locator(".ant-select").filter({ hasText: "Select primary model" });
await primarySelect.click();
// FallbackGroupConfig.tsx renders both fields as searchable comboboxes: they
// open on click, typing filters the listbox, and the option has to be picked
// explicitly. Verify each selection landed by watching the dialog's own state
// transition (the tab title updates to the picked primary; the fallback chain
// list populates) rather than by asserting on the popup, which is portaled
// out of the dialog.
await modal.getByRole("combobox", { name: /Primary Model/ }).click();
await page.keyboard.type(PRIMARY);
await page.keyboard.press("Enter");
await page.getByRole("option", { name: PRIMARY, exact: true }).click();
await expect(modal.getByRole("tab", { name: PRIMARY })).toBeVisible({
timeout: 10_000,
});
const fallbackSelect = modal.locator(".ant-select").filter({ hasText: "Select fallback models" });
await fallbackSelect.click();
await modal.getByRole("combobox", { name: /Select fallback models/ }).click();
await page.keyboard.type(FALLBACK);
await page.keyboard.press("Enter");
await page.keyboard.press("Escape");
await page.getByRole("option", { name: FALLBACK, exact: true }).click();
// The Fallback Chain helper text reads "(N/10 used)"; once it ticks to 1 the
// selection has been recorded.
await expect(modal.getByText("(1/10 used)")).toBeVisible({

View file

@ -61,7 +61,7 @@ test.describe("Team Admin", () => {
await page.getByRole("tab", { name: "Members" }).click();
await page.getByRole("button", { name: /Add Member/i }).click();
const modal = page.locator(".ant-modal:visible");
const modal = page.getByRole("dialog", { name: "Add Team Member" });
await expect(modal).toBeVisible({ timeout: 5_000 });
// Use a dedicated invitee user so this doesn't race with the proxy-admin
@ -144,9 +144,10 @@ test.describe("Team Admin", () => {
await page.keyboard.type(E2E_TEAM_CRUD_ALIAS);
await page.locator('[data-slot="combobox-content"]:visible').getByText(E2E_TEAM_CRUD_ALIAS).first().click();
// Models — pick "All Team Models"
await page.locator(".ant-select-selection-overflow").click();
await page.locator(".ant-select-dropdown:visible").getByText("All Team Models").click();
// Models — pick "All Team Models". The popup is portaled to the body, so
// scope the option lookup to the page.
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "All Team Models", exact: true }).click();
await page.keyboard.press("Escape");
const generate = await captureRequestBody(page, { method: "POST", urlIncludes: "/key/generate" }, async () => {

View file

@ -22,7 +22,8 @@ async function openUsage(page: PlaywrightPage): Promise<Locator> {
const card = topKeysCard(page);
await expect(card).toBeVisible({ timeout: 30_000 });
// Widen past the default top-5 so other keys in the database cannot crowd this one out.
await card.locator(".ant-segmented-item").filter({ hasText: /^50$/ }).click();
// The radio itself is sr-only and its label covers it, so click the label.
await card.getByRole("radiogroup", { name: "Number of top keys to show" }).getByText("50", { exact: true }).click();
return card;
}

View file

@ -11,7 +11,7 @@ test.skip("Internal Users Search", () => {
await tab.click();
await expect(page.locator("tbody tr").first()).toBeVisible();
await expect(page.locator(".ant-skeleton")).toHaveCount(0);
await expect(page.locator('[data-slot="skeleton"]')).toHaveCount(0);
}
test("can search users by email", async ({ page }) => {

View file

@ -13,7 +13,7 @@ test.skip("Internal Users Page", () => {
const firstRow = page.locator("tbody tr").first();
await expect(firstRow).toBeVisible();
await expect(page.locator(".ant-skeleton")).toHaveCount(0);
await expect(page.locator('[data-slot="skeleton"]')).toHaveCount(0);
}
test("renders internal users table correctly", async ({ page }) => {

View file

@ -0,0 +1,163 @@
"""Tests for the shared PTU rules: which deployments accrue flat cost, and what that zeroes."""
import os
from datetime import datetime, timezone
from unittest.mock import patch
import pytest
from litellm.litellm_core_utils.ptu_pricing import (
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
PTU_ZEROED_PRICING_FIELDS,
PTU_ZEROED_TABLE_FIELDS,
SEARCH_CONTEXT_SIZES,
ptu_terms,
zeroed_ptu_pricing,
)
from litellm.types.router import ModelInfo
_VALID = {
"team_id": "team-alpha",
"ptu_count": 100,
"cost_per_ptu_per_hour": 0.02,
"ptu_effective_from": "2026-01-01T00:00:00Z",
}
def _with_flag(model_info, declared=None, enabled=True):
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True" if enabled else ""}, clear=False):
return zeroed_ptu_pricing(model_info, declared or {})
def test_a_complete_reservation_is_accepted():
terms = ptu_terms(_VALID)
assert terms is not None
assert terms.team_id == "team-alpha"
assert terms.ptu_count == 100
assert terms.effective_from == datetime(2026, 1, 1, tzinfo=timezone.utc)
assert terms.effective_to is None
@pytest.mark.parametrize(
"override",
[
{"team_id": None},
{"team_id": ""},
{"ptu_count": None},
{"cost_per_ptu_per_hour": None},
{"ptu_count": 0},
{"ptu_count": -1},
{"ptu_count": ModelInfo.MAX_PTU_COUNT + 1},
{"cost_per_ptu_per_hour": -0.01},
{"cost_per_ptu_per_hour": ModelInfo.MAX_COST_PER_PTU_PER_HOUR + 1},
{"ptu_count": "not-a-number"},
{"ptu_effective_from": None},
{"ptu_effective_from": "not-a-date"},
{"ptu_effective_to": "not-a-date"},
{"ptu_effective_to": "2025-01-01T00:00:00Z"},
{"ptu_effective_to": "2026-01-01T00:00:00Z"},
],
ids=[
"no team",
"blank team",
"no count",
"no rate",
"zero count",
"negative count",
"count over the cap",
"negative rate",
"rate over the cap",
"count not a number",
"no start",
"unparseable start",
"unparseable end",
"end before start",
"end equal to start",
],
)
def test_an_incomplete_reservation_accrues_nothing(override):
"""Anything the rollup declines to charge must also decline to be zeroed, or the
deployment serves its traffic for free with nothing charged in its place."""
assert ptu_terms({**_VALID, **override}) is None
assert _with_flag({**_VALID, **override}) is None
def test_a_naive_start_is_read_as_utc():
"""config.yaml is hand-typed, and pydantic hands back a naive datetime for a date with
no offset."""
terms = ptu_terms({**_VALID, "ptu_effective_from": datetime(2026, 5, 1, 12, 0)})
assert terms is not None
assert terms.effective_from == datetime(2026, 5, 1, 12, 0, tzinfo=timezone.utc)
def test_an_offset_start_is_converted_rather_than_relabelled():
terms = ptu_terms({**_VALID, "ptu_effective_from": "2026-05-01T12:00:00-05:00"})
assert terms is not None
assert terms.effective_from == datetime(2026, 5, 1, 17, 0, tzinfo=timezone.utc)
def test_nothing_is_zeroed_while_the_feature_is_off():
"""No flat cost accrues with the flag off, so zeroing would serve the traffic free."""
assert _with_flag(_VALID, enabled=False) is None
def test_the_standing_rates_are_all_zeroed():
override = _with_flag(_VALID)
assert override is not None
assert [field for field in PTU_ZEROED_PRICING_FIELDS if override[field] != 0.0] == []
def test_tiered_pricing_is_emptied_rather_than_zeroed():
"""A tier outranks the flat rates written beside it, so a zero there would leave the
cost map's tiers billing the traffic the reserved capacity already covers."""
override = _with_flag(_VALID, declared={"tiered_pricing": [{"range": [0, 1000], "input_cost_per_token": 0.003}]})
assert override is not None
for field in PTU_EMPTIED_PRICING_FIELDS:
assert override[field] == ()
def test_the_search_context_table_is_zeroed_in_place_on_every_deployment():
"""An absent table means the provider's own default rather than free, so it is written
even when the deployment never declared one."""
override = _with_flag(_VALID)
assert override is not None
for field in PTU_ZEROED_TABLE_FIELDS:
assert dict(override[field]) == dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0)
def test_a_declared_table_does_not_become_a_scalar():
"""Zeroing it as a plain 0.0 would leave the provider's reader without a table to
consult, which is the same as absent."""
override = _with_flag(_VALID, declared={"search_context_cost_per_query": {"search_context_size_medium": 0.05}})
assert override is not None
assert dict(override["search_context_cost_per_query"]) == dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0)
def test_a_rate_the_deployment_declares_itself_is_zeroed_too():
"""The standing set covers the mirrored rates. Anything else the operator wrote would
otherwise survive and bill the traffic the hourly charge already paid for."""
extra = "input_cost_per_token_above_200k_tokens"
assert extra in CUSTOM_PRICING_FIELDS
assert extra not in PTU_ZEROED_PRICING_FIELDS
override = _with_flag(_VALID, declared={extra: 9e-06})
assert override is not None
assert override[extra] == 0.0
def test_a_setting_that_is_not_a_charge_is_left_alone():
"""CustomPricingLiteLLMParams also carries configuration, and zeroing one of those
would break the deployment rather than stop a charge."""
override = _with_flag(_VALID, declared={"output_vector_size": 1536})
assert override is not None
assert "output_vector_size" not in override

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
import sys
@ -20,9 +21,11 @@ class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator):
def __init__(self, litellm_logging_obj: LiteLLMLoggingObj, request_body: dict):
super().__init__(litellm_logging_obj=litellm_logging_obj, request_body=request_body)
self.logged_chunks: list = []
self.logging_call_count: int = 0
async def _handle_streaming_logging(self, collected_chunks):
self.logged_chunks = list(collected_chunks)
self.logging_call_count += 1
def _make_logging_obj(test_name: str) -> LiteLLMLoggingObj:
@ -233,6 +236,70 @@ async def test_async_sse_wrapper_excludes_synthetic_error_event_from_logged_chun
assert not any(chunk.startswith(b"event: error\n") for chunk in iterator.logged_chunks)
async def _events_then_hang(events):
for event in events:
yield event
await asyncio.Event().wait()
@pytest.mark.asyncio
async def test_async_sse_wrapper_logs_partial_chunks_on_client_disconnect():
"""
Regression test for LIT-5839: a client disconnect tears the generator
down with GeneratorExit at the yield, which used to skip the post-loop
logging dispatch entirely, so the partial output tokens the provider
already generated (and billed) never reached spend tracking.
"""
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_disconnect_logs_partial_chunks"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
assert iterator.logging_call_count == 0
await wrapped.aclose()
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == streamed
@pytest.mark.asyncio
async def test_async_sse_wrapper_logs_partial_chunks_on_cancellation():
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_cancellation_logs_partial_chunks"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
consume_task = asyncio.ensure_future(wrapped.__anext__())
await asyncio.sleep(0.01)
consume_task.cancel()
with pytest.raises(asyncio.CancelledError):
await consume_task
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == streamed
@pytest.mark.asyncio
async def test_async_sse_wrapper_skips_logging_on_disconnect_before_first_chunk():
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_disconnect_before_first_chunk"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(()))
consume_task = asyncio.ensure_future(wrapped.__anext__())
await asyncio.sleep(0.01)
consume_task.cancel()
with pytest.raises(asyncio.CancelledError):
await consume_task
assert iterator.logging_call_count == 0
def test_incomplete_stream_error_sse_event_is_valid_anthropic_error():
event = _incomplete_stream_error_sse_event().decode()
lines = event.split("\n")

View file

@ -2948,3 +2948,38 @@ def test_bedrock_invoke_messages_allows_converted_websearch_function_tool():
headers={},
)
assert result["tools"][0]["name"] == "litellm_web_search"
@pytest.mark.asyncio
async def test_bedrock_sse_wrapper_dispatches_logging_on_client_disconnect():
"""
Regression test for LIT-5839: closing the outer bedrock_sse_wrapper
mid-stream (what the proxy does on a client disconnect) must close the
inner async_sse_wrapper deterministically so the partial-stream logging
fires. `completion_start_time` is only stamped on the logging object by
that dispatch, so it observing a value proves the whole chain ran.
"""
cfg = AmazonAnthropicClaudeMessagesConfig()
async def _hanging_stream():
yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 25, "output_tokens": 1}}}
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}
await asyncio.Event().wait()
logging_obj = LiteLLMLoggingObj(
model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="chat",
start_time=datetime.now(),
litellm_call_id="test_bedrock_sse_wrapper_disconnect_logging",
function_id="test_bedrock_sse_wrapper_disconnect_logging",
)
wrapped = cfg.bedrock_sse_wrapper(_hanging_stream(), litellm_logging_obj=logging_obj, request_body={})
await wrapped.__anext__()
await wrapped.__anext__()
assert logging_obj.completion_start_time is None
await wrapped.aclose()
assert logging_obj.completion_start_time is not None

View file

@ -25,6 +25,11 @@ from litellm.proxy._types import (
from litellm.proxy.common_utils.key_rotation_manager import KeyRotationManager
@pytest.fixture
def disable_audit_logging_for_mocked_key(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.store_audit_logs", False)
class TestKeyRotationManagerPassesKeyAlias:
"""
Regression tests to ensure KeyRotationManager passes key_alias
@ -155,7 +160,10 @@ class TestKeyRotationSecretNamingStability:
"""
@pytest.mark.asyncio
async def test_rotation_hook_uses_initial_secret_name_fallback(self):
async def test_rotation_hook_uses_initial_secret_name_fallback(
self,
disable_audit_logging_for_mocked_key,
):
"""
GIVEN: A key WITHOUT an alias (has an initial_secret_name based on token ID)
WHEN: The key is rotated
@ -206,7 +214,10 @@ class TestKeyRotationSecretNamingStability:
), f"Secret name drift! Expected {initial_secret_name}, got {call_kwargs['new_secret_name']}. This causes secret sprawl."
@pytest.mark.asyncio
async def test_rotation_hook_pre_rotation_alias_consistency(self):
async def test_rotation_hook_pre_rotation_alias_consistency(
self,
disable_audit_logging_for_mocked_key,
):
"""
GIVEN: A key WITH an alias
WHEN: The key is rotated

View file

@ -0,0 +1,441 @@
"""
LIT-5273: enqueued-token accounting for batch submissions.
Covers the ``BatchEnqueuedTokenStore`` (reserve / refund / reservation
records), the metadata-driven scope resolution, and the batch-id and
response-shape helpers the v3 limiter's post-call hooks rely on.
"""
import base64
import socket
import uuid
from collections.abc import Mapping, Sequence
from types import MappingProxyType, SimpleNamespace
from typing import Final
import pytest
from litellm.caching.caching import DualCache
from litellm.constants import BATCH_ENQUEUED_TOKEN_TTL_SECONDS
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenOverLimit,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
BatchEnqueuedTokenStore,
batch_response_view,
canonical_provider_batch_id,
resolve_batch_enqueued_token_scopes,
)
from litellm.proxy.utils import InternalUsageCache
def _in_memory_store() -> BatchEnqueuedTokenStore:
return BatchEnqueuedTokenStore(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60)))
def _scope(limit: int, key: str = "api_key") -> BatchEnqueuedTokenScope:
return BatchEnqueuedTokenScope(key=key, value=f"{key}-{uuid.uuid4().hex}", limit=limit)
def test_scope_resolution_reads_key_and_team_metadata():
user = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"batch_enqueued_token_limit": 100},
team_id="team-1",
team_metadata={"batch_enqueued_token_limit": "150"},
)
scopes = resolve_batch_enqueued_token_scopes(user)
assert scopes == (
BatchEnqueuedTokenScope(key="api_key", value="hashed-key", limit=100),
BatchEnqueuedTokenScope(key="team", value="team-1", limit=150),
)
def test_scope_resolution_returns_empty_without_opt_in():
assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k")) == ()
assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k", metadata={}, team_metadata=None)) == ()
@pytest.mark.parametrize("bad_value", ["not-a-number", 0, -5, None, [1000]])
def test_scope_resolution_ignores_invalid_limits(bad_value):
user = UserAPIKeyAuth(api_key="k", metadata={"batch_enqueued_token_limit": bad_value})
assert resolve_batch_enqueued_token_scopes(user) == ()
def test_scope_resolution_skips_team_scope_without_team_id():
user = UserAPIKeyAuth(api_key="k", team_metadata={"batch_enqueued_token_limit": 100})
assert resolve_batch_enqueued_token_scopes(user) == ()
@pytest.mark.asyncio
async def test_reserve_rejects_once_allowance_is_exhausted():
store = _in_memory_store()
scope = _scope(limit=100)
first = await store.reserve(tokens=80, scopes=(scope,))
assert isinstance(first, BatchEnqueuedTokenReservation)
second = await store.reserve(tokens=30, scopes=(scope,))
assert second == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=80)
third = await store.reserve(tokens=20, scopes=(scope,))
assert isinstance(third, BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_reserve_is_all_or_nothing_across_scopes():
store = _in_memory_store()
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=50, key="team")
over = await store.reserve(tokens=60, scopes=(key_scope, team_scope))
assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
exact_fit = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(exact_fit, BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_refund_restores_allowance_and_never_goes_negative():
store = _in_memory_store()
scope = _scope(limit=100)
reservation = await store.reserve(tokens=30, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
await store.refund(reservation)
await store.refund(reservation)
refill = await store.reserve(tokens=100, scopes=(scope,))
assert isinstance(refill, BatchEnqueuedTokenReservation)
assert isinstance(await store.reserve(tokens=1, scopes=(scope,)), BatchEnqueuedTokenOverLimit)
@pytest.mark.asyncio
async def test_reservation_record_roundtrip_pops_exactly_once():
store = _in_memory_store()
scope = _scope(limit=100)
reservation = await store.reserve(tokens=40, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
await store.save_reservation("batch_abc", reservation)
assert await store.pop_reservation("batch_abc") == reservation
assert await store.pop_reservation("batch_abc") is None
assert await store.pop_reservation("batch_never_saved") is None
@pytest.mark.asyncio
async def test_zero_token_reserve_charges_nothing():
store = _in_memory_store()
scope = _scope(limit=100)
empty = await store.reserve(tokens=0, scopes=(scope,))
assert empty == BatchEnqueuedTokenReservation(tokens=0, scopes=(scope,))
full = await store.reserve(tokens=100, scopes=(scope,))
assert isinstance(full, BatchEnqueuedTokenReservation)
class _SingleKeyRedisFake:
"""Emulates the Redis script path one single-key call at a time, recording every call."""
def __init__(
self,
fail_reserve_keys: frozenset[str] = frozenset(),
fail_refund_keys: frozenset[str] = frozenset(),
fail_save_keys: frozenset[str] = frozenset(),
raise_after_landing_save_keys: frozenset[str] = frozenset(),
) -> None:
self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = ()
self.save_ttls: tuple[int, ...] = ()
self.counters: Mapping[str, int] = MappingProxyType({})
self.records: Mapping[str, str] = MappingProxyType({})
self.fail_reserve_keys = fail_reserve_keys
self.fail_refund_keys = fail_refund_keys
self.fail_save_keys = fail_save_keys
self.raise_after_landing_save_keys = raise_after_landing_save_keys
def async_register_script(self, script: str):
kind: Final = (
"reserve"
if "INCRBY" in script
else "refund" if "DECRBY" in script else "pop" if "GET" in script else "save"
)
async def run(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object:
self.script_calls = (*self.script_calls, (kind, tuple(keys)))
return self._run(kind, tuple(keys), tuple(args))
return run
def _run(self, kind: str, keys: tuple[str, ...], args: tuple[str | bytes | int | float, ...]) -> object:
if kind == "reserve":
if keys[0] in self.fail_reserve_keys:
raise ConnectionError(f"simulated redis failure for {keys[0]}")
amount, limit = int(args[0]), int(args[2])
current: Final = self.counters.get(keys[0], 0)
if current + amount > limit:
return (0, current)
self.counters = MappingProxyType({**self.counters, keys[0]: current + amount})
return (1, current + amount)
if kind == "refund":
if keys[0] in self.fail_refund_keys:
raise ConnectionError(f"simulated redis failure for {keys[0]}")
remaining: Final = self.counters.get(keys[0], 0) - int(args[0])
self.counters = MappingProxyType(
{key: value for key, value in self.counters.items() if key != keys[0]}
if remaining <= 0
else {**self.counters, keys[0]: remaining}
)
return 1
if kind == "save":
if keys[0] in self.fail_save_keys:
raise ConnectionError(f"simulated redis failure for {keys[0]}")
self.records = MappingProxyType({**self.records, keys[0]: str(args[0])})
self.save_ttls = (*self.save_ttls, int(args[1]))
if keys[0] in self.raise_after_landing_save_keys:
raise TimeoutError(f"simulated redis timeout after landing for {keys[0]}")
return 1
if kind == "pop":
popped: Final = self.records.get(keys[0])
if popped:
self.records = MappingProxyType({**self.records, keys[0]: ""})
return popped
raise AssertionError(f"unexpected {kind} script call for keys {keys}")
@pytest.mark.asyncio
async def test_redis_reserve_issues_single_key_calls_and_rolls_back_on_over_limit():
fake = _SingleKeyRedisFake()
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=50, key="team")
over = await store.reserve(tokens=60, scopes=(key_scope, team_scope))
assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund")
assert not fake.counters
fits = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(fits, BatchEnqueuedTokenReservation)
await store.refund(fits)
assert not fake.counters
assert all(len(keys) == 1 for _, keys in fake.script_calls)
@pytest.mark.asyncio
async def test_partial_redis_reserve_failure_rolls_back_and_grants_in_memory():
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=50, key="team")
fake = _SingleKeyRedisFake(fail_reserve_keys=frozenset({f"batch_enqueued_tokens:team:{team_scope.value}"}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
outcome = await store.reserve(tokens=10, scopes=(key_scope, team_scope))
assert isinstance(outcome, BatchEnqueuedTokenReservation)
assert outcome.backend == "memory"
assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund")
assert not fake.counters
await store.refund(outcome)
assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund")
refilled = await store.reserve(tokens=50, scopes=(team_scope,))
assert isinstance(refilled, BatchEnqueuedTokenReservation)
assert refilled.backend == "memory"
@pytest.mark.asyncio
async def test_over_limit_verdict_survives_a_failing_rollback():
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=5, key="team")
key_counter: Final = f"batch_enqueued_tokens:api_key:{key_scope.value}"
fake = _SingleKeyRedisFake(fail_refund_keys=frozenset({key_counter}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
outcome = await store.reserve(tokens=10, scopes=(key_scope, team_scope))
assert outcome == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
assert fake.counters == {key_counter: 10}
@pytest.mark.asyncio
async def test_pop_falls_back_to_local_record_when_redis_pop_finds_nothing():
scope = _scope(limit=100)
record_key: Final = "batch_enqueued_token_reservation:batch_local_record"
fake = _SingleKeyRedisFake(fail_save_keys=frozenset({record_key}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
await store.save_reservation("batch_local_record", reservation)
assert not fake.records
popped = await store.pop_reservation("batch_local_record")
assert popped == reservation
await store.refund(popped)
assert not fake.counters
assert await store.pop_reservation("batch_local_record") is None
@pytest.mark.asyncio
async def test_local_ghost_left_by_landed_save_never_refunds_twice():
scope = _scope(limit=100)
record_key: Final = "batch_enqueued_token_reservation:batch_ghost"
fake = _SingleKeyRedisFake(raise_after_landing_save_keys=frozenset({record_key}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
await store.save_reservation("batch_ghost", reservation)
assert fake.records[record_key]
first = await store.pop_reservation("batch_ghost")
assert first == reservation
await store.refund(first)
assert not fake.counters
assert fake.records[record_key] == ""
assert await store.pop_reservation("batch_ghost") is None
assert (
await store.internal_usage_cache.async_get_cache(
key=record_key, litellm_parent_otel_span=None, local_only=True
)
is None
)
assert await store.pop_reservation("batch_ghost") is None
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
assert fake.counters[f"batch_enqueued_tokens:{scope.key}:{scope.value}"] == 100
@pytest.mark.asyncio
async def test_record_ttl_shrinks_by_elapsed_time_so_stale_records_never_outlive_their_counters():
scope = _scope(limit=100)
fake = _SingleKeyRedisFake()
ticks = iter((1_000.0, 1_030.5))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)),
monotonic=lambda: next(ticks),
)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.reserved_at_monotonic == 1_000.0
await store.save_reservation("batch_ttl_clamp", reservation)
assert fake.save_ttls == (BATCH_ENQUEUED_TOKEN_TTL_SECONDS - 31,)
assert await store.pop_reservation("batch_ttl_clamp") == reservation
@pytest.mark.asyncio
async def test_memory_refund_skips_reservations_granted_by_another_worker():
store = _in_memory_store()
scope = _scope(limit=100)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.backend == "memory"
assert reservation.owner
foreign: Final = BatchEnqueuedTokenReservation(
tokens=60, scopes=reservation.scopes, backend="memory", owner="another-worker"
)
await store.refund(foreign)
assert await store.reserve(tokens=50, scopes=(scope,)) == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=60)
await store.refund(reservation)
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_failed_redis_refund_leaves_local_counters_untouched():
scope = _scope(limit=100)
counter_key: Final = f"batch_enqueued_tokens:api_key:{scope.value}"
fake = _SingleKeyRedisFake(fail_refund_keys=frozenset({counter_key}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.backend == "redis"
store.internal_usage_cache.dual_cache.in_memory_cache.set_cache(key=counter_key, value=45)
await store.refund(reservation)
assert store.internal_usage_cache.dual_cache.in_memory_cache.get_cache(key=counter_key) == 45
assert fake.counters == {counter_key: 60}
@pytest.mark.asyncio
async def test_pop_reservation_defaults_legacy_records_to_redis_backend():
store = _in_memory_store()
legacy = '{"tokens": 5, "scopes": [{"key": "api_key", "value": "k", "limit": 10}]}'
store.internal_usage_cache.dual_cache.in_memory_cache.set_cache(
key="batch_enqueued_token_reservation:batch_legacy", value=legacy
)
popped = await store.pop_reservation("batch_legacy")
assert popped == BatchEnqueuedTokenReservation(
tokens=5, scopes=(BatchEnqueuedTokenScope(key="api_key", value="k", limit=10),), backend="redis"
)
def test_canonical_provider_batch_id_passes_raw_ids_through():
assert canonical_provider_batch_id("batch_abc123") == "batch_abc123"
def test_canonical_provider_batch_id_decodes_unified_batch_ids():
unified = "litellm_proxy;model_id:m-1;llm_batch_id:batch_prov_9;llm_output_file_id:file-9"
encoded = base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=")
assert canonical_provider_batch_id(encoded) == "batch_prov_9"
def test_canonical_provider_batch_id_decodes_model_embedded_ids():
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
encoded = encode_file_id_with_model(file_id="batch_prov_7", model="my-alias", id_type="batch")
assert canonical_provider_batch_id(encoded) == "batch_prov_7"
def test_batch_response_view_accepts_batch_objects_only():
batch = SimpleNamespace(id="batch_1", status="completed", object="batch")
view = batch_response_view(batch)
assert view is not None and view.id == "batch_1" and view.status == "completed"
assert batch_response_view({"id": "chatcmpl-1", "object": "chat.completion"}) is None
assert batch_response_view(None) is None
assert batch_response_view("batch_1") is None
def _local_redis_port() -> int | None:
for port in (6379,):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.2)
if sock.connect_ex(("127.0.0.1", port)) == 0:
return port
return None
@pytest.mark.asyncio
@pytest.mark.skipif(_local_redis_port() is None, reason="requires a local Redis on 6379 for the Lua script path")
async def test_redis_lua_path_full_lifecycle():
from litellm.caching.redis_cache import RedisCache
port = _local_redis_port()
redis_cache = RedisCache(host="127.0.0.1", port=port)
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache, default_in_memory_ttl=60))
)
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=50, key="team")
over = await store.reserve(tokens=60, scopes=(key_scope, team_scope))
assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
reservation = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert isinstance(await store.reserve(tokens=1, scopes=(key_scope, team_scope)), BatchEnqueuedTokenOverLimit)
batch_id = f"batch_{uuid.uuid4().hex}"
await store.save_reservation(batch_id, reservation)
popped = await store.pop_reservation(batch_id)
assert popped == reservation
assert await store.pop_reservation(batch_id) is None
await store.refund(popped)
refill = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(refill, BatchEnqueuedTokenReservation)
await store.refund(refill)

View file

@ -2094,3 +2094,195 @@ def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, exp
}
assert rate_limiter._estimate_entry_output_tokens(entry, None) == expected
# ---------------------------------------------------------------------------
# LIT-5273: enqueued-token limits govern batch submission when opted in
# ---------------------------------------------------------------------------
def _enqueued_rate_limiter():
from litellm import DualCache
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.utils import InternalUsageCache
local_cache = DualCache(default_in_memory_ttl=60)
internal_usage_cache = InternalUsageCache(local_cache)
parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache)
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=internal_usage_cache,
parallel_request_limiter=parallel_request_limiter,
)
return rate_limiter, local_cache
_ENQUEUED_BATCH_FILE_CONTENT = (
b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n'
b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n'
b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n'
)
def _enqueued_batch_patches():
mock_content = MagicMock()
mock_content.content = _ENQUEUED_BATCH_FILE_CONTENT
afile_content_mock = AsyncMock(return_value=mock_content)
return afile_content_mock, (
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={"custom_llm_provider": "openai"},
),
)
@pytest.mark.asyncio
async def test_enqueued_limit_accepts_batch_over_per_minute_limits():
"""The headline LIT-5273 behavior: a key that opted into an enqueued-token
allowance submits a batch whose row count and token count both exceed its
per-minute RPM/TPM limits, and the batch is accepted (repeatedly) because
only the enqueued allowance governs. Without the opt-in the same key is
rejected on RPM before the batch reaches the provider."""
from litellm.proxy.hooks.parallel_request_limiter_v3 import get_request_stash
rate_limiter, local_cache = _enqueued_rate_limiter()
afile_content_mock, patches = _enqueued_batch_patches()
legacy_user = UserAPIKeyAuth(api_key="sk-legacy-rpm", models=["*"], rpm_limit=1, tpm_limit=10)
opted_in_user = UserAPIKeyAuth(
api_key="sk-enqueued-rpm",
models=["*"],
rpm_limit=1,
tpm_limit=10,
metadata={"batch_enqueued_token_limit": 100000},
)
with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock):
with pytest.raises(HTTPException) as legacy_exc:
await rate_limiter.async_pre_call_hook(
user_api_key_dict=legacy_user,
cache=local_cache,
data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"},
call_type="acreate_batch",
)
assert legacy_exc.value.status_code == 429
first_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"}
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=opted_in_user,
cache=local_cache,
data=first_data,
call_type="acreate_batch",
)
assert result is first_data
stash = get_request_stash()
assert stash is not None and stash.batch_enqueued_reservation is not None
assert stash.batch_enqueued_reservation.tokens == first_data["_batch_token_count"] > 0
second = await rate_limiter.async_pre_call_hook(
user_api_key_dict=opted_in_user,
cache=local_cache,
data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"},
call_type="acreate_batch",
)
assert second is not None
@pytest.mark.asyncio
async def test_enqueued_limit_rejects_when_allowance_is_exhausted():
"""Submissions are rejected pre-provider once the enqueued allowance can't
fit the batch, even for a key with no per-minute limits at all (which
previously skipped batch rate limiting entirely)."""
rate_limiter, local_cache = _enqueued_rate_limiter()
afile_content_mock, patches = _enqueued_batch_patches()
sizing_user = UserAPIKeyAuth(
api_key="sk-enqueued-sizing", models=["*"], metadata={"batch_enqueued_token_limit": 1000000}
)
with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock):
sizing_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"}
await rate_limiter.async_pre_call_hook(
user_api_key_dict=sizing_user,
cache=local_cache,
data=sizing_data,
call_type="acreate_batch",
)
batch_tokens = sizing_data["_batch_token_count"]
assert batch_tokens > 0
capped_user = UserAPIKeyAuth(
api_key="sk-enqueued-capped",
models=["*"],
metadata={"batch_enqueued_token_limit": batch_tokens + batch_tokens // 2},
)
await rate_limiter.async_pre_call_hook(
user_api_key_dict=capped_user,
cache=local_cache,
data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"},
call_type="acreate_batch",
)
with pytest.raises(HTTPException) as exc:
await rate_limiter.async_pre_call_hook(
user_api_key_dict=capped_user,
cache=local_cache,
data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"},
call_type="acreate_batch",
)
assert exc.value.status_code == 429
assert "Batch enqueued token limit exceeded for api_key" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_enqueued_team_limit_applies_to_batch_submission():
rate_limiter, local_cache = _enqueued_rate_limiter()
afile_content_mock, patches = _enqueued_batch_patches()
team_user = UserAPIKeyAuth(
api_key="sk-enqueued-team-key",
models=["*"],
team_id="team-enqueued-batch",
team_metadata={"batch_enqueued_token_limit": 10},
)
with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock):
with pytest.raises(HTTPException) as exc:
await rate_limiter.async_pre_call_hook(
user_api_key_dict=team_user,
cache=local_cache,
data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"},
call_type="acreate_batch",
)
assert exc.value.status_code == 429
assert "Batch enqueued token limit exceeded for team: team-enqueued-batch" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_disable_flag_still_skips_batch_processing_with_enqueued_limits():
rate_limiter, local_cache = _enqueued_rate_limiter()
afile_content_mock, _ = _enqueued_batch_patches()
opted_in_user = UserAPIKeyAuth(
api_key="sk-enqueued-disabled",
models=["*"],
metadata={"batch_enqueued_token_limit": 10},
)
with (
patch("litellm.proxy.proxy_server.general_settings", {"disable_batch_input_file_rate_limiting": True}),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch("litellm.afile_content", new=afile_content_mock),
):
data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"}
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=opted_in_user,
cache=local_cache,
data=data,
call_type="acreate_batch",
)
assert result is data
afile_content_mock.assert_not_awaited()

View file

@ -4,6 +4,7 @@ Tests for KeyManagementEventHooks.
Validates that email and secret manager operations are independent and non-blocking.
"""
import asyncio
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -155,6 +156,44 @@ class TestKeyManagementEventHooksIndependentOperations:
assert email_called["called"] is True
@pytest.mark.parametrize(
("premium_user", "expected_audit_log_calls"),
((True, 1), (False, 0)),
)
@pytest.mark.asyncio
async def test_key_generated_audit_log_uses_license_default(
monkeypatch: pytest.MonkeyPatch,
premium_user: bool,
expected_audit_log_calls: int,
):
from litellm.proxy._types import GenerateKeyRequest, GenerateKeyResponse, UserAPIKeyAuth
monkeypatch.setattr("litellm.store_audit_logs", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", premium_user)
monkeypatch.delenv("LITELLM_STORE_AUDIT_LOGS", raising=False)
response = GenerateKeyResponse(key="sk-test-key", token_id="token-123")
with (
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new_callable=AsyncMock,
) as mock_create_audit_log,
patch.object(
KeyManagementEventHooks,
"_store_virtual_key_in_secret_manager",
new_callable=AsyncMock,
),
):
await KeyManagementEventHooks.async_key_generated_hook(
data=GenerateKeyRequest(),
response=response,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin-key", user_id="admin"),
)
await asyncio.sleep(0.01)
assert mock_create_audit_log.await_count == expected_audit_log_calls
class TestRotateVirtualKeyInSecretManager:
"""Tests for _rotate_virtual_key_in_secret_manager with team_id support."""

View file

@ -5881,3 +5881,150 @@ def test_parallel_request_slot_ttl_env_override():
check=True,
)
assert output.stdout.strip() == "300"
# ---------------------------------------------------------------------------
# LIT-5273: batch enqueued-token reservations in the post-call hooks
# ---------------------------------------------------------------------------
def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler:
return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60)))
def _batch_response(batch_id: str, status: str):
from types import SimpleNamespace
return SimpleNamespace(id=batch_id, status=status, object="batch")
@pytest.mark.asyncio
async def test_success_hook_persists_batch_enqueued_reservation_and_refunds_on_completion():
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenOverLimit,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
)
handler = _enqueued_test_handler()
store = handler.batch_enqueued_token_store
scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-enqueued-key", limit=100)
user = UserAPIKeyAuth(api_key="hashed-enqueued-key")
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
get_or_create_request_stash().batch_enqueued_reservation = reservation
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "validating")
)
assert get_request_stash().batch_enqueued_reservation is None
assert isinstance(await store.reserve(tokens=50, scopes=(scope,)), BatchEnqueuedTokenOverLimit)
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed")
)
refill = await store.reserve(tokens=40, scopes=(scope,))
assert isinstance(refill, BatchEnqueuedTokenReservation)
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed")
)
assert isinstance(await store.reserve(tokens=70, scopes=(scope,)), BatchEnqueuedTokenOverLimit)
@pytest.mark.asyncio
async def test_success_hook_refunds_batch_enqueued_reservation_on_cancellation():
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
)
handler = _enqueued_test_handler()
store = handler.batch_enqueued_token_store
scope = BatchEnqueuedTokenScope(key="team", value="team-enqueued", limit=100)
user = UserAPIKeyAuth(api_key="hashed-enqueued-key", team_id="team-enqueued")
reservation = await store.reserve(tokens=90, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
get_or_create_request_stash().batch_enqueued_reservation = reservation
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "validating")
)
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "cancelling")
)
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_success_hook_refunds_on_provider_cased_terminal_status():
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenOverLimit,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
)
handler = _enqueued_test_handler()
store = handler.batch_enqueued_token_store
scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-enqueued-key", limit=100)
user = UserAPIKeyAuth(api_key="hashed-enqueued-key")
reservation = await store.reserve(tokens=90, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
get_or_create_request_stash().batch_enqueued_reservation = reservation
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_cased", "InProgress")
)
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenOverLimit)
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=_batch_response("batch_enq_cased", "Completed")
)
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_failure_hook_refunds_stashed_batch_enqueued_reservation():
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
)
handler = _enqueued_test_handler()
store = handler.batch_enqueued_token_store
scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-failing-key", limit=100)
user = UserAPIKeyAuth(api_key="hashed-failing-key")
reservation = await store.reserve(tokens=80, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
get_or_create_request_stash().batch_enqueued_reservation = reservation
await handler.async_post_call_failure_hook(
request_data={}, original_exception=Exception("guardrail rejected"), user_api_key_dict=user
)
assert get_request_stash().batch_enqueued_reservation is None
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_success_hook_leaves_stash_untouched_for_non_batch_responses():
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
)
handler = _enqueued_test_handler()
store = handler.batch_enqueued_token_store
scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-chat-key", limit=100)
user = UserAPIKeyAuth(api_key="hashed-chat-key")
reservation = await store.reserve(tokens=10, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
get_or_create_request_stash().batch_enqueued_reservation = reservation
await handler.async_post_call_success_hook(
data={}, user_api_key_dict=user, response=ModelResponse(usage=Usage(total_tokens=5))
)
assert get_request_stash().batch_enqueued_reservation == reservation

View file

@ -212,8 +212,9 @@ async def test_v1_key_generation_sends_email_when_send_invite_email_true():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.slack_alerting_instance = mock_slack_alerting
with patch.object(
KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email
with (
patch("litellm.store_audit_logs", False),
patch.object(KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email),
):
with patch(
"litellm.logging_callback_manager.get_custom_loggers_for_type",
@ -257,8 +258,9 @@ async def test_v1_key_generation_no_email_when_send_invite_email_false():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.slack_alerting_instance = mock_slack_alerting
with patch.object(
KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email
with (
patch("litellm.store_audit_logs", False),
patch.object(KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email),
):
with patch(
"litellm.logging_callback_manager.get_custom_loggers_for_type",

View file

@ -15519,6 +15519,207 @@ async def test_regenerate_key_output_token_estimate_lowered_rejected_for_non_adm
assert "Only proxy admins can set" in str(exc.value.detail)
_BATCH_LIMIT = "batch_enqueued_token_limit"
@pytest.mark.parametrize(
"label, request_body, existing_metadata, allowed",
[
("set on a key with none stored", {"metadata": {_BATCH_LIMIT: 50000}}, None, False),
("raised above the stored limit", {"metadata": {_BATCH_LIMIT: 200000}}, {_BATCH_LIMIT: 100000}, False),
("cleared by replacing the blob", {"metadata": {}}, {_BATCH_LIMIT: 100000}, False),
("resent unchanged", {"metadata": {_BATCH_LIMIT: 100000}}, {_BATCH_LIMIT: 100000}, True),
("left untouched", {}, {_BATCH_LIMIT: 100000}, True),
],
)
def test_batch_enqueued_token_limit_admin_gate_matrix(label, request_body, existing_metadata, allowed):
"""A non-admin may only leave a key's stored batch enqueued-token limit as it is.
When set, the limit replaces the standard RPM/TPM checks for batch
submissions, so a key holder writing it would pick their own batch quota.
Resending the stored value is what the edit form produces on every save
and has to stay allowed.
"""
from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
)
def _call(caller):
enforce_batch_enqueued_token_limit_is_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_batch_enqueued_token_limit_rejected_for_non_admin():
"""A non-admin self-minting a key with the limit would replace the standard
batch RPM/TPM checks with a cap of their own choosing."""
with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()):
with pytest.raises(HTTPException) as exc:
await _common_key_generation_helper(
data=GenerateKeyRequest(metadata={_BATCH_LIMIT: 100000}, rpm_limit=2),
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_update_key_batch_enqueued_token_limit_raised_rejected_for_non_admin(monkeypatch):
"""/key/update is reachable by the key's own holder, so the gate has to
fire inside the update path itself rather than only at generation."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
token = "d1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
_wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_BATCH_LIMIT: 100000}))
mock_request = MagicMock()
mock_request.query_params = {}
with pytest.raises(ProxyException) as exc:
await update_key_fn(
request=mock_request,
data=UpdateKeyRequest(key=token, metadata={_BATCH_LIMIT: 10**12}),
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_batch_enqueued_token_limit_unchanged_allows_non_admin_edit(monkeypatch):
"""The edit form resends every field it renders, so gating on presence
would 403 a key owner renaming a key that carries an admin-set limit."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
token = "e1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
_wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_BATCH_LIMIT: 100000}))
mock_request = MagicMock()
mock_request.query_params = {}
result = await update_key_fn(
request=mock_request,
data=UpdateKeyRequest(key=token, key_alias="my-alias", metadata={_BATCH_LIMIT: 100000}),
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_batch_enqueued_token_limit_rejected_for_non_admin():
"""/key/regenerate runs the request body through prepare_key_update_data
exactly as an update does, so it is a third write path into the field."""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
token = "f1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
key_in_db = LiteLLM_VerificationToken(
token=token,
user_id="internal_user",
metadata={_BATCH_LIMIT: 100000},
)
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", metadata={_BATCH_LIMIT: 10**12}),
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)
@pytest.mark.asyncio
async def test_bulk_key_update_batch_enqueued_token_limit_rejected_for_non_admin():
"""Bulk team-key updates run through _process_single_key_update, not
/key/update's validator, so the gate must also live on that path."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_process_single_key_update,
)
token = "a2b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
existing = _estimate_key_row(token, {_BATCH_LIMIT: 100000})
with pytest.raises(HTTPException) as exc:
await _process_single_key_update(
update_key_request=UpdateKeyRequest(key=token, metadata={_BATCH_LIMIT: 10**12}),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-internal",
user_id="internal_user",
),
litellm_changed_by=None,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
llm_router=None,
existing_key_row=existing,
)
assert exc.value.status_code == 403
assert "Only proxy admins can set" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_stamps_settings_updated_at():
"""Regenerate rewrites the key's config, so it must move settings_updated_at."""

View file

@ -91,6 +91,8 @@ mock_prisma_client = MagicMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_teamtable = MagicMock()
mock_prisma_client.db.litellm_teamtable.update = AsyncMock()
mock_prisma_client.db.litellm_auditlog = MagicMock()
mock_prisma_client.db.litellm_auditlog.create = AsyncMock()
# Fixture to provide the mock prisma client
@ -103,6 +105,11 @@ def mock_db_client():
mock_prisma_client.reset_mock()
@pytest.fixture
def disable_audit_logging_for_mocked_team(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.store_audit_logs", False)
# Fixture to provide a mock admin user auth object
@pytest.fixture
def mock_admin_auth():
@ -2060,7 +2067,9 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name):
@pytest.mark.asyncio
async def test_update_team_team_member_budget_not_passed_to_db():
async def test_update_team_team_member_budget_not_passed_to_db(
disable_audit_logging_for_mocked_team,
):
"""
Test that 'team_member_budget' is never passed to prisma_client.db.litellm_teamtable.update
regardless of whether the value is set or None.
@ -2498,7 +2507,9 @@ async def test_upsert_team_member_budget_table_no_existing_budget():
@pytest.mark.asyncio
async def test_update_team_with_team_member_budget_duration():
async def test_update_team_with_team_member_budget_duration(
disable_audit_logging_for_mocked_team,
):
"""
Test that team/update endpoint properly handles team_member_budget_duration.
"""
@ -5171,7 +5182,9 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin():
@pytest.mark.asyncio
async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin():
async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin(
disable_audit_logging_for_mocked_team,
):
"""
Test that a proxy admin CAN raise a standalone team's budget on /team/update.
@ -5325,7 +5338,9 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin():
@pytest.mark.asyncio
async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed():
async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed(
disable_audit_logging_for_mocked_team,
):
"""
When a team currently has NO cap (max_budget=None / unlimited), a team admin
setting a finite max_budget is a RESTRICTION, not a raise, and is
@ -5407,7 +5422,9 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed():
@pytest.mark.asyncio
async def test_update_team_standalone_unchanged_budget_allowed():
async def test_update_team_standalone_unchanged_budget_allowed(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team does NOT compare against the
caller's personal max_budget when the budget is unchanged.
@ -5508,7 +5525,9 @@ async def test_update_team_standalone_unchanged_budget_allowed():
@pytest.mark.asyncio
async def test_update_team_standalone_lower_budget_allowed():
async def test_update_team_standalone_lower_budget_allowed(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team allows lowering the budget
below the team's current value even when the new value still exceeds the
@ -5691,7 +5710,9 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit():
@pytest.mark.asyncio
async def test_update_team_standalone_models_not_gated_by_user_limit():
async def test_update_team_standalone_models_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team does NOT gate the team's models
by the caller's personal allowed models.
@ -5775,7 +5796,9 @@ async def test_update_team_standalone_models_not_gated_by_user_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_budget_bypasses_user_limit():
async def test_update_team_org_scoped_budget_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team does NOT validate budget against user's personal max_budget.
@ -5890,7 +5913,9 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_models_bypasses_user_limit():
async def test_update_team_org_scoped_models_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team does NOT validate models against user's personal models.
@ -6080,7 +6105,9 @@ async def test_update_team_org_scoped_models_not_in_org_models():
@pytest.mark.asyncio
async def test_update_team_org_scoped_models_with_all_proxy_models():
async def test_update_team_org_scoped_models_with_all_proxy_models(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team succeeds when organization has 'all-proxy-models'.
@ -6196,7 +6223,9 @@ async def test_update_team_org_scoped_models_with_all_proxy_models():
@pytest.mark.asyncio
async def test_update_team_tpm_limit_not_gated_by_user_limit():
async def test_update_team_tpm_limit_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update does NOT gate the team's tpm_limit by the caller's
personal tpm_limit.
@ -6279,7 +6308,9 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit():
@pytest.mark.asyncio
async def test_update_team_rpm_limit_not_gated_by_user_limit():
async def test_update_team_rpm_limit_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update does NOT gate the team's rpm_limit by the caller's
personal rpm_limit.
@ -6795,7 +6826,9 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit():
async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team bypasses user's TPM/RPM limits.
@ -6905,7 +6938,9 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit():
@pytest.mark.asyncio
async def test_update_team_guardrails_with_org_id():
async def test_update_team_guardrails_with_org_id(
disable_audit_logging_for_mocked_team,
):
"""
Test that updating team guardrails works when team has an organization_id.
The fix ensures 'teams' field is included when fetching organization data.
@ -7242,7 +7277,10 @@ async def test_persist_deleted_team_records():
@pytest.mark.asyncio
async def test_delete_team_persists_deleted_teams(monkeypatch):
async def test_delete_team_persists_deleted_teams(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
from litellm.proxy._types import DeleteTeamRequest
mock_prisma_client = AsyncMock()
@ -7325,7 +7363,10 @@ async def test_delete_team_persists_deleted_teams(monkeypatch):
@pytest.mark.asyncio
async def test_delete_team_sweeps_references_outside_members_with_roles(monkeypatch):
async def test_delete_team_sweeps_references_outside_members_with_roles(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Regression pin for LIT-5511: a deleted team stayed visible on user records.
@ -7431,7 +7472,10 @@ async def test_delete_team_sweeps_references_outside_members_with_roles(monkeypa
@pytest.mark.asyncio
async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(monkeypatch):
async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
A virtual key scoped to the team is deleted from the db with the team, but auth resolves a
cached key object without re-reading the team, so leaving the cache entry behind lets that key
@ -7493,7 +7537,10 @@ async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(monkeypa
@pytest.mark.asyncio
async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cache(monkeypatch):
async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cache(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
The reconcile sweep runs after the team row is committed deleted. If it ran before cache
eviction, a sweep failure would return an error with the team gone from the db but still
@ -7555,7 +7602,10 @@ async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cac
@pytest.mark.asyncio
async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(monkeypatch):
async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Evicting locally only reaches the worker that handled the delete. Without the broadcast, every
other worker keeps serving the deleted team, and the deleted team's keys, out of its own
@ -7619,7 +7669,10 @@ async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(monkey
@pytest.mark.asyncio
async def test_delete_team_survives_a_failing_cache_backend(monkeypatch):
async def test_delete_team_survives_a_failing_cache_backend(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Cache eviction runs after the reference sweep has already committed, so a cache backend that
is unreachable must not abort the delete. If it did, `/team/delete` would fail with the team
@ -8088,6 +8141,7 @@ async def test_update_team_soft_budget_validation(
expected_soft_budget,
expected_max_budget,
error_message,
disable_audit_logging_for_mocked_team,
):
"""
Test soft_budget validation in /team/update endpoint.
@ -8498,7 +8552,11 @@ async def test_get_team_daily_activity_member_without_permission_filters_by_keys
@pytest.mark.asyncio
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
async def test_update_team_with_router_settings(
mock_db_client,
mock_admin_auth,
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update correctly handles router_settings by:
1. Accepting router_settings as a dict parameter
@ -11594,7 +11652,9 @@ async def test_update_team_output_token_estimate_lowered_rejected_for_team_admin
@pytest.mark.asyncio
async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit():
async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit(
disable_audit_logging_for_mocked_team,
):
"""The team settings form resends every field it renders, so gating on
presence would break a team admin editing an unrelated setting."""
import contextlib
@ -11649,6 +11709,63 @@ async def test_new_team_output_token_estimate_rejected_for_non_admin():
assert "on a team" in str(exc.value.message)
_TEAM_BATCH_LIMIT = "batch_enqueued_token_limit"
@pytest.mark.asyncio
async def test_update_team_batch_enqueued_token_limit_raised_rejected_for_team_admin():
"""_verify_team_access admits a team admin, so the gate has to fire inside
update_team itself to keep the team's batch quota admin-owned."""
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_BATCH_LIMIT: 100000})
with pytest.raises(ProxyException) as exc:
await update_team(
data=UpdateTeamRequest(team_id="test_team_id", metadata={_TEAM_BATCH_LIMIT: 10**12}),
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_new_team_batch_enqueued_token_limit_rejected_for_non_admin():
"""/team/new is the other write path into the same stored metadata."""
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", metadata={_TEAM_BATCH_LIMIT: 100000}),
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)
@pytest.mark.asyncio
async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_client):
"""The aggregated endpoint must apply the same non-admin key scoping as the
@ -11957,7 +12074,9 @@ class _FakeMirrorDb:
@pytest.mark.asyncio
async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directions():
async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directions(
disable_audit_logging_for_mocked_team,
):
"""
A team-side edit of `access_group_ids` must be mirrored onto every affected access
group's `assigned_team_ids`, in one transaction, in both directions.
@ -12113,7 +12232,9 @@ async def test_sync_reads_the_committed_team_row_rather_than_the_callers_snapsho
@pytest.mark.asyncio
async def test_new_team_and_delete_team_both_drive_the_mirror():
async def test_new_team_and_delete_team_both_drive_the_mirror(
disable_audit_logging_for_mocked_team,
):
"""Every writer of `team.access_group_ids` has to reach the mirror, not just update.
These pin the wiring on the other two paths; the mirror's own behavior is covered above.

View file

@ -19,6 +19,7 @@ from litellm.proxy.management_helpers.audit_logs import (
_build_audit_log_payload,
_dispatch_audit_log_to_callbacks,
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.types.utils import StandardAuditLogPayload
@ -49,6 +50,34 @@ def _make_audit_log(
)
@pytest.mark.parametrize(
("premium_user", "configured_value", "environment_value", "expected"),
(
(True, None, None, True),
(True, False, None, False),
(True, None, "false", False),
(False, None, None, False),
(False, True, None, True),
(True, True, "false", True),
),
)
def test_is_audit_logging_enabled_precedence(
monkeypatch: pytest.MonkeyPatch,
premium_user: bool,
configured_value: bool | None,
environment_value: str | None,
expected: bool,
):
monkeypatch.setattr(litellm, "store_audit_logs", configured_value)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", premium_user)
if environment_value is None:
monkeypatch.delenv("LITELLM_STORE_AUDIT_LOGS", raising=False)
else:
monkeypatch.setenv("LITELLM_STORE_AUDIT_LOGS", environment_value)
assert is_audit_logging_enabled() is expected
class TestBuildAuditLogPayload:
def test_builds_correct_payload(self):
audit_log = _make_audit_log()
@ -185,12 +214,14 @@ class TestCreateAuditLogForUpdateWithCallbacks:
with (
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.store_audit_logs", True),
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
):
audit_log = _make_audit_log()
await create_audit_log_for_update(audit_log)
await asyncio.sleep(0.1)
mock_logger.async_log_audit_log_event.assert_not_called()
mock_prisma.db.litellm_auditlog.create.assert_not_called()
@pytest.mark.asyncio
async def test_no_dispatch_when_store_audit_logs_false(self):

View file

@ -210,8 +210,10 @@ async def test_rollup_prunes_stale_row_when_config_is_gone():
where = table.delete_many.await_args.kwargs["where"]
assert where["date"] == DAY.isoformat()
assert where["api_key"] == PTU_SENTINEL_API_KEY
# the row is garbage because this run did not refresh it, not because of a key list
# the row is garbage because this run did not refresh it, and it is reachable at all
# because the run scanned the deployment it belongs to
assert "lt" in where["updated_at"]
assert "model" not in where, "a database-only run has no reason to bound the sweep"
@pytest.mark.asyncio
@ -705,13 +707,20 @@ class _FakeSentinelTable:
async def delete_many(self, where):
self.delete_many_calls.append(where)
cutoff = where["updated_at"]["lt"]
# honouring "model" matters: a fake that ignored an unknown clause would delete
# the row the prune-scoping test exists to protect and still report a pass
allowed = where.get("model", {}).get("in")
doomed = [
k
for k, v in self.rows.items()
if k[1] == where["date"] and k[2] == where["api_key"] and v["updated_at"] < cutoff
if k[1] == where["date"]
and k[2] == where["api_key"]
and v["updated_at"] < cutoff
and (allowed is None or k[3] in allowed)
]
for k in doomed:
del self.rows[k]
return len(doomed)
async def find_many(self, where=None):
"""Read back sentinel rows the way prisma would, honouring api_key and a date range."""
@ -785,11 +794,11 @@ async def test_an_older_run_cannot_delete_a_newer_runs_row():
@pytest.mark.asyncio
async def test_a_later_clean_run_clears_the_row_the_race_left_behind():
"""The race can leave a charge for a since-removed deployment in place for a day; the
next run, seeing only the current config, must sweep it."""
"""The race can leave a charge for a no-longer-priced deployment in place for a day;
the next run, seeing only the current config, must sweep it."""
table = _FakeSentinelTable()
ptu = {"ptu_count": 10, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
stale_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-removed")
stale_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-retired")
table.rows[stale_key] = {
"ptu_flat_cost": 480.0,
"model_group": "retired",
@ -797,7 +806,14 @@ async def test_a_later_clean_run_clears_the_row_the_race_left_behind():
}
await run_ptu_flat_cost_rollup(
_prisma_for([_model_row(model_id="dep-live", model_info=ptu)], table), target_date=DAY
_prisma_for(
[
_model_row(model_id="dep-live", model_info=ptu),
_model_row(model_id="dep-retired", model_info={"team_id": "t"}),
],
table,
),
target_date=DAY,
)
assert stale_key not in table.rows
@ -1715,16 +1731,19 @@ async def test_a_run_holding_the_lock_still_prunes():
"""Losing the sweep entirely would leave stale charges forever, so the guarded path,
which is the normal one, keeps it."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-gone", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
table.seed("t", DAY, "dep-unpriced", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
prisma = _prisma_for(
[_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
[
_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}),
_model_row(model_id="dep-unpriced", model_info={"team_id": "t"}),
],
table,
)
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert table.delete_many_calls != []
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-gone") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-unpriced") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-live") in table.rows
@ -1737,7 +1756,7 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
just_written = datetime.now(timezone.utc) - timedelta(seconds=30)
table.seed("t", DAY, "dep-concurrent", 480.0, updated_at=just_written)
table.seed("t", DAY, "dep-stale", 480.0, updated_at=datetime.now(timezone.utc) - timedelta(hours=6))
prisma = _prisma_for([], table)
prisma = _prisma_for([_model_row(model_id="dep-concurrent"), _model_row(model_id="dep-stale")], table)
await run_ptu_flat_cost_rollup(prisma, target_date=DAY)
@ -1747,6 +1766,127 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-stale") not in table.rows
@pytest.mark.asyncio
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypatch):
"""Staleness alone stops being evidence once two hosts hold different configuration: a
row this run never considered belongs to a deployment another host is pricing from its
own file, and sweeping it drops that charge."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-elsewhere", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
entry = _router_entry(model_id="cfg-here", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-elsewhere") in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-here") in table.rows
assert table.delete_many_calls[-1]["model"]["in"] == ("cfg-here",)
@pytest.mark.asyncio
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(monkeypatch):
"""The accepted cost of bounding the prune, driven through the sequence that produces
it: charge the day while the deployment exists, remove it, run the day again. Nothing
scans it now, so nothing may judge its row, and the amount it was billed stands."""
table = _FakeSentinelTable()
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
live_row = _model_row(model_id="dep-live", model_info=ptu)
doomed_row = _model_row(model_id="dep-doomed", model_info=ptu)
charged_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-doomed")
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
)
await run_scheduled_ptu_rollup(
_prisma_for([live_row, doomed_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
billed = table.rows[charged_key]["ptu_flat_cost"]
table.rows[charged_key]["updated_at"] = datetime(2020, 1, 1, tzinfo=timezone.utc)
await run_scheduled_ptu_rollup(
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
assert table.rows[charged_key]["ptu_flat_cost"] == billed
assert "dep-doomed" not in table.delete_many_calls[-1]["model"]["in"]
@pytest.mark.asyncio
async def test_a_database_only_run_sweeps_exactly_as_it_did_before():
"""The bound exists for charges another host declares. A deployment nobody declares any
more still has its leftover row swept, which is what the table-only sweep always did."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-gone", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
prisma = _prisma_for(
[_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
table,
)
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-gone") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-live") in table.rows
@pytest.mark.asyncio
async def test_every_deployment_that_prices_is_inside_the_set_that_bounds_the_prune():
"""The bound has to be a superset of what the same run wrote, or a run's own charge
could fall outside its own delete filter and never be reconciled."""
table = _FakeSentinelTable()
prisma = _prisma_for(
[
_model_row(model_id="dep-a", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}),
_model_row(model_id="dep-b", model_info={"ptu_count": 9, "cost_per_ptu_per_hour": 1.0, "team_id": "u"}),
_model_row(model_id="dep-unpriced", model_info={"team_id": "t"}),
],
table,
)
loaded = await ptu_rollup._load_ptu_models(prisma)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
assert loaded.scanned_ids == {"dep-a", "dep-b", "dep-unpriced"}
@pytest.mark.asyncio
async def test_a_priced_deployment_is_in_the_bound_even_with_an_id_the_scan_skips():
"""The bound is built by construction rather than by coincidence. The row scan drops a
falsy id while the parser still prices one, and a charge outside its own run's delete
filter could never be reconciled by any later run."""
prisma = _prisma_for(
[_model_row(model_id="", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
_FakeSentinelTable(),
)
loaded = await ptu_rollup._load_ptu_models(prisma)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
@pytest.mark.asyncio
async def test_the_prune_splits_the_id_set_across_statements(monkeypatch):
"""Every id is one bind variable and the server refuses a statement carrying more than
32767, so a proxy with that many deployments would fail the prune outright, and with it
the rest of the scheduled run."""
monkeypatch.setattr(ptu_rollup, "_PRUNE_ID_CHUNK_SIZE", 2)
table = _FakeSentinelTable()
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
deployments = [_model_row(model_id=f"dep-{n}", model_info=ptu) for n in range(4)]
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu)))
)
table.seed("t", DAY, "dep-3", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
await run_scheduled_ptu_rollup(
_prisma_for(deployments, table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
chunks = [call["model"]["in"] for call in table.delete_many_calls]
assert len(chunks) == 3
assert all(len(chunk) <= 2 for chunk in chunks)
assert sorted(i for chunk in chunks for i in chunk) == [f"dep-{n}" for n in range(5)]
@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
@ -1760,3 +1900,185 @@ async def test_scheduled_rollup_writes_nothing_when_ptu_attribution_is_disabled(
assert result is None
assert table.rows == {}
assert table.upsert_keys == []
# --- config.yaml deployments reach the rollup through the router ----------------
def _router_holding(*entries):
"""A stand-in for the proxy's router, carrying whatever model_list is passed."""
return types.SimpleNamespace(model_list=list(entries))
@pytest.mark.asyncio
async def test_a_config_declared_deployment_is_priced(monkeypatch):
"""The whole point. A PTU deployment the proxy only knows from config.yaml is not in
LiteLLM_ProxyModelTable, so a DB-only scan bills the provider's reservation to nobody."""
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert [(m.model_id, m.model_name, m.team_id) for m in loaded.models] == [("cfg-1", "gpt-4o-ptu", "t")]
assert "cfg-1" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_a_database_backed_router_entry_is_not_counted_twice(monkeypatch):
"""Every deployment loaded from the table is also in the router, flagged db_model. Pricing
both copies would write two charges for one reservation."""
row = _model_row(model_id="db-1", model_info=dict(_VALID_PTU))
mirrored = _router_entry(model_id="db-1", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(mirrored))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once(monkeypatch):
"""db_model is data the router carries rather than something this module controls, so the
id anti-join is what actually maps onto the failure: two charges under one id."""
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(unflagged))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["both-1"]
@pytest.mark.asyncio
async def test_a_client_credential_clone_is_not_priced(monkeypatch):
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(source, clone))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["cfg-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced(monkeypatch):
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert loaded.models == ()
assert "cfg-plain" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_no_router_in_the_process_prices_the_database_alone(monkeypatch):
"""The rollup is importable and callable outside a running proxy."""
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: None)
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable())
)
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_is_charged_end_to_end(monkeypatch):
"""Through the scheduled entry point, so the charge lands in a sentinel row rather than
stopping at the loader."""
table = _FakeSentinelTable()
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-1") in table.rows
@pytest.mark.asyncio
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config(monkeypatch):
"""The reconcile can leave a deployment on the router after its row is gone. The id
anti-join cannot see that one, so the flag is what keeps it from being priced as though
config.yaml had declared it."""
stale = _router_entry(model_id="db-gone", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(stale))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert loaded.models == ()
def test_the_router_lookup_reads_the_proxys_own_global():
"""Every other config test replaces this helper, so without one test driving the real
body a typo in the module path or the attribute name leaves the whole feature dead in
production with the suite still green."""
import sys
import types as _types
assert ptu_rollup._running_router() is None or "litellm.proxy.proxy_server" in sys.modules
sentinel = object()
stub = _types.SimpleNamespace(llm_router=sentinel)
real = sys.modules.get("litellm.proxy.proxy_server")
sys.modules["litellm.proxy.proxy_server"] = stub
try:
assert ptu_rollup._running_router() is sentinel
del stub.llm_router
assert ptu_rollup._running_router() is None
finally:
if real is None:
del sys.modules["litellm.proxy.proxy_server"]
else:
sys.modules["litellm.proxy.proxy_server"] = real
def test_the_router_lookup_returns_none_outside_a_proxy():
import sys
real = sys.modules.pop("litellm.proxy.proxy_server", None)
try:
assert ptu_rollup._running_router() is None
finally:
if real is not None:
sys.modules["litellm.proxy.proxy_server"] = real
@pytest.mark.parametrize("chunk", [None, ("dep-a", "dep-b")], ids=["unbounded", "bounded"])
def test_the_prune_filter_is_a_plain_dict(chunk):
"""The query builder serialises the mapping it is handed and rejects a read-only view of
one, which the in-memory table in these tests accepts happily. Only a live run caught it."""
predicate = ptu_rollup._prune_filter(date_str=DAY.isoformat(), cutoff=datetime.now(timezone.utc), chunk=chunk)
assert type(predicate) is dict
assert type(predicate["updated_at"]) is dict
if chunk is None:
assert "model" not in predicate
else:
assert type(predicate["model"]) is dict
assert predicate["model"]["in"] == chunk
@pytest.mark.asyncio
async def test_the_catch_up_pass_reaches_a_config_declared_deployment(monkeypatch):
"""The catch-up shares the loader, so config deployments join it without being wired in.
That is what prices the elapsed days of a reservation declared before today."""
table = _FakeSentinelTable()
now = datetime.now(timezone.utc)
started = (now - timedelta(days=3)).strftime("%Y-%m-%dT00:00:00Z")
entry = _router_entry(
model_id="cfg-back",
model_info={"ptu_count": 100, "cost_per_ptu_per_hour": 0.02, "team_id": "t", "ptu_effective_from": started},
)
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True))
charged = sorted(day for (_, day, _, model) in table.rows if model == "cfg-back")
yesterday = (now.date() - timedelta(days=1)).isoformat()
assert len(charged) == 3, charged
assert charged[-1] == yesterday
assert all(row["ptu_flat_cost"] == pytest.approx(48.0) for row in table.rows.values())

View file

@ -3164,3 +3164,80 @@ def test_batch_cost_row_id_is_stable_across_repeated_accounting():
]
assert ids[0] == ids[1] == "batch_same_batch_cost"
def _make_failed_request_standard_logging_payload() -> StandardLoggingPayload:
base: Final = _make_standard_logging_payload_with_usage_object(usage_object={})
return cast(
StandardLoggingPayload,
{
**base,
"status": "failure",
"call_type": "aresponses",
"model_id": "mid-123",
"model_group": "group-x",
"api_base": "https://api.openai.com/v1/responses",
"custom_llm_provider": "openai",
},
)
def test_get_logging_payload_failed_request_falls_back_to_standard_logging_payload():
"""Failed-request kwargs from the proxy failure hook carry no deployment info
(LIT-5795), so the attribution columns must come from the failure-time
standard_logging_object."""
payload = get_logging_payload(
kwargs={
"model": "group-x",
"litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}},
"standard_logging_object": _make_failed_request_standard_logging_payload(),
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == "mid-123"
assert payload["model_group"] == "group-x"
assert payload["api_base"] == "https://api.openai.com/v1/responses"
assert payload["custom_llm_provider"] == "openai"
def test_get_logging_payload_request_kwargs_win_over_standard_logging_payload():
payload = get_logging_payload(
kwargs={
"model": "group-y",
"custom_llm_provider": "anthropic",
"litellm_params": {
"api_base": "https://kwargs.example.com",
"metadata": {
"user_api_key": "test-key",
"model_group": "kwargs-group",
"model_info": {"id": "kwargs-mid"},
},
},
"standard_logging_object": _make_failed_request_standard_logging_payload(),
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == "kwargs-mid"
assert payload["model_group"] == "kwargs-group"
assert payload["api_base"] == "https://kwargs.example.com"
assert payload["custom_llm_provider"] == "anthropic"
def test_get_logging_payload_failed_request_without_standard_logging_payload_leaves_fields_empty():
payload = get_logging_payload(
kwargs={
"model": "group-x",
"litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}},
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == ""
assert payload["model_group"] == ""
assert payload["api_base"] == ""
assert payload["custom_llm_provider"] == ""

View file

@ -468,6 +468,23 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert request_data["response_cost"] == 3.5e-05
assert "litellm_logging_obj" not in request_data
@pytest.mark.asyncio
async def test_recovered_usage_without_cost_clobbers_client_cost_with_zero(self):
from litellm.types.utils import Usage
recovered_usage = Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31)
logging_obj = MagicMock()
logging_obj.model_call_details = {"combined_usage_object": recovered_usage}
request_data = {
"litellm_logging_obj": logging_obj,
"response_cost": 999.0,
"metadata": {},
}
await self._run(request_data)
assert request_data["combined_usage_object"] is recovered_usage
assert request_data["response_cost"] == 0.0
@pytest.mark.asyncio
async def test_no_recovered_usage_is_noop(self):
logging_obj = MagicMock()
@ -478,6 +495,111 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert "response_cost" not in request_data
class TestPostCallFailureHookLiftsStandardLoggingObject:
"""Failure callbacks read standard_logging_object from request_data, but
post_call_failure_hook pops litellm_logging_obj before they run. The hook
must lift the logging obj's standard_logging_object onto request_data so
failed-request spend logs keep deployment attribution (LIT-5795).
"""
async def _run(self, request_data):
from unittest.mock import AsyncMock, patch
from litellm.proxy._types import UserAPIKeyAuth
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data=request_data,
original_exception=Exception("boom"),
user_api_key_dict=UserAPIKeyAuth(),
)
@pytest.mark.asyncio
async def test_lifts_standard_logging_object(self):
sl_object = {"model_id": "mid-123", "model_group": "group-x"}
logging_obj = MagicMock()
logging_obj.model_call_details = {"standard_logging_object": sl_object}
request_data = {"litellm_logging_obj": logging_obj, "metadata": {}}
await self._run(request_data)
assert request_data["standard_logging_object"] is sl_object
assert "litellm_logging_obj" not in request_data
@pytest.mark.asyncio
async def test_logging_obj_value_overwrites_preexisting_key(self):
authoritative = {"model_id": "from-logging-obj"}
logging_obj = MagicMock()
logging_obj.model_call_details = {"standard_logging_object": authoritative}
request_data = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": {"model_id": "client-injected"},
"metadata": {},
}
await self._run(request_data)
assert request_data["standard_logging_object"] is authoritative
@pytest.mark.asyncio
async def test_client_supplied_key_is_stripped_when_logging_obj_supplies_none(self):
spoofed = {"model_id": "client-injected"}
request_data = {"standard_logging_object": spoofed, "metadata": {}}
await self._run(request_data)
assert "standard_logging_object" not in request_data
logging_obj = MagicMock()
logging_obj.model_call_details = {}
request_data_with_obj = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": spoofed,
"metadata": {},
}
await self._run(request_data_with_obj)
assert "standard_logging_object" not in request_data_with_obj
@pytest.mark.asyncio
async def test_pass_through_failure_never_relifts_client_supplied_key(self):
from datetime import datetime
from unittest.mock import AsyncMock, patch
from fastapi import HTTPException
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy._types import UserAPIKeyAuth
logging_obj = Logging(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="test-call-id",
function_id="test-function-id",
)
request_data = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": {"model_id": "client-injected"},
"metadata": {},
}
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data=request_data,
original_exception=HTTPException(status_code=401, detail="unauthorized"),
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"),
)
assert "standard_logging_object" not in request_data
assert "standard_logging_object" not in logging_obj.model_call_details
@pytest.mark.asyncio
async def test_no_standard_logging_object_is_noop(self):
logging_obj = MagicMock()
logging_obj.model_call_details = {}
request_data = {"litellm_logging_obj": logging_obj, "metadata": {}}
await self._run(request_data)
assert "standard_logging_object" not in request_data
class TestPostCallFailureHookEstimatesDispatchedInputTokens:
"""A non-stream request that failed after dispatch (timeout, provider
error) consumed provider-billed input tokens but recovered no usage.

View file

@ -2166,6 +2166,38 @@ class TestRouterPreRoutingAliasOverrides:
assert request_kwargs["drop_params"] is True
assert request_kwargs["cache_control_injection_points"] == [{"location": "message", "role": "system"}]
@pytest.mark.asyncio
async def test_tier_litellm_params_are_applied_before_deployment_selection(self):
router = Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"tiers": {
"SIMPLE": {
"model_name": "gpt-4o-mini",
"litellm_params": {"reasoning_effort": "xhigh"},
}
}
},
},
},
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}},
]
)
request_kwargs: Dict = {"reasoning_effort": "low"}
deployment = await router.async_get_available_deployment(
model="smart-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hi"}],
)
assert deployment["model_name"] == "gpt-4o-mini"
assert request_kwargs["reasoning_effort"] == "xhigh"
@pytest.mark.asyncio
async def test_alias_custom_pricing_is_not_applied_to_request_kwargs(self):
"""Custom pricing on the alias prices the alias, not the tier deployment
@ -3851,7 +3883,7 @@ class TestSessionAffinity:
cache.async_set_cache.assert_called_once()
call_kwargs = cache.async_set_cache.call_args.kwargs
assert call_kwargs["ttl"] == 120
assert call_kwargs["value"] == "gpt-4o-mini"
assert call_kwargs["value"] == {"model": "gpt-4o-mini", "tier": "SIMPLE"}
@pytest.mark.asyncio
async def test_ttl_refreshed_on_cache_hit(self, mock_router_instance, basic_config):
@ -3875,7 +3907,7 @@ class TestSessionAffinity:
assert result.model == "o1-preview"
cache.async_set_cache.assert_called_once()
call_kwargs = cache.async_set_cache.call_args.kwargs
assert call_kwargs["value"] == "o1-preview"
assert call_kwargs["value"] == {"model": "o1-preview", "tier": "REASONING"}
assert call_kwargs["ttl"] == 90
@pytest.mark.asyncio
@ -5545,12 +5577,14 @@ class TestRedactedLoggingDropsPromptText:
"tier_boundaries": {"simple_medium": 0.15, "medium_complex": 0.35, "complex_reasoning": 0.6},
"classifier_model": "claude-haiku",
"escalated": True,
"tier_litellm_params": {"reasoning_effort": "xhigh"},
"signals": ["code (python)"],
"matched_keyword": "deploy to k8s",
"escalation_keyword": "LITELLM ESCALATE",
}
kept = Router._redact_prompt_text_if_needed(request_kwargs={}, routing_decision=full)
assert set(full) - set(kept) == {"signals", "matched_keyword", "escalation_keyword"}
assert kept["tier_litellm_params"] == {"reasoning_effort": "xhigh"}
@pytest.mark.asyncio
async def test_redaction_via_request_header_is_honored(self):
@ -7273,8 +7307,6 @@ class TestClassificationRubrics:
},
)
assert config.classifier_llm_config.system_prompt == "Grade the data sensitivity of the request."
def _custom_tier_config(**overrides) -> Dict:
"""A valid operator-defined tier set: two built-in names plus one custom tier."""
return {
@ -8123,3 +8155,243 @@ class TestPlanModeTierFloor:
assert result.model == "gpt-4o"
assert result.routing_decision is not None
assert result.routing_decision["tier"] == "MEDIUM"
def test_tier_model_params_are_normalized_without_changing_model_pools():
config = ComplexityRouterConfig(
tiers={
"SIMPLE": "mini",
"REASONING": [
{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}},
"abc",
],
}
)
assert config.tiers == {"SIMPLE": "mini", "REASONING": ["opus", "abc"]}
assert config.tier_model_configs["REASONING"][0].litellm_params == {"reasoning_effort": "xhigh"}
rebuilt = ComplexityRouterConfig.model_validate(config.model_dump())
assert rebuilt.tier_model_configs["REASONING"][0].litellm_params == {"reasoning_effort": "xhigh"}
def test_tier_model_params_accept_a_single_object():
config = ComplexityRouterConfig(
tiers={"REASONING": {"model_name": "opus", "litellm_params": {"thinking": {"type": "enabled"}}}}
)
assert config.tiers == {"REASONING": "opus"}
assert config.tier_model_configs["REASONING"][0].model_name == "opus"
@pytest.mark.parametrize(
"tiers",
[
{"REASONING": [{"litellm_params": {"reasoning_effort": "xhigh"}}]},
],
)
def test_tier_model_params_reject_malformed_entries(tiers):
with pytest.raises(ValidationError):
ComplexityRouterConfig(tiers=tiers)
def test_tier_model_params_reject_duplicate_models():
with pytest.raises(ValidationError, match="duplicate model_name"):
ComplexityRouterConfig(
tiers={
"REASONING": [
{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}},
{"model_name": "opus", "litellm_params": {"reasoning_effort": "low"}},
]
}
)
def test_non_adaptive_empty_tier_pool_remains_valid():
config = ComplexityRouterConfig(tiers={"SIMPLE": []})
assert config.tiers == {"SIMPLE": []}
def test_adaptive_empty_tier_pool_is_rejected():
with pytest.raises(ValidationError, match="adaptive=True"):
ComplexityRouterConfig(adaptive=True, tiers={"SIMPLE": []})
def test_tier_model_params_are_used_by_pools_and_savings_baseline(mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": "mini",
"REASONING": [{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}}, "abc"],
}
},
)
assert router._tier_pools() == {"SIMPLE": ["mini"], "REASONING": ["opus", "abc"]}
assert router._hardest_tier_models() == ("opus", "abc")
assert router._litellm_params_for_model(ComplexityTier.REASONING, "opus") == {"reasoning_effort": "xhigh"}
@pytest.mark.asyncio
async def test_tier_model_params_reach_the_hook_response_and_override_client_values(mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"REASONING": {
"model_name": "opus",
"litellm_params": {"reasoning_effort": "xhigh", "max_tokens": 512},
}
},
"keyword_tier_rules": [{"keywords": ["reason carefully"], "tier": "REASONING"}],
},
)
request_kwargs = {"reasoning_effort": "low", "metadata": {}}
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "reason carefully about this"}],
)
assert response is not None
assert response.litellm_params == {"reasoning_effort": "xhigh", "max_tokens": 512}
assert response.routing_decision is not None
assert response.routing_decision["tier_litellm_params"] == response.litellm_params
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["classification", "keyword", "session"])
async def test_tier_params_mask_credentials_in_routing_decision(route, mock_router_instance):
params = {"reasoning_effort": "xhigh", "api_key": "secret-tier-key"}
config = {
"tiers": {
tier.value: {"model_name": "opus", "litellm_params": params}
for tier in ComplexityTier
},
"keyword_tier_rules": [{"keywords": ["reason carefully"], "tier": "REASONING"}]
if route == "keyword"
else None,
"session_affinity": route == "session",
}
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=config,
)
request_kwargs = {"metadata": {"session_id": "masked-params-session"}}
if route == "session":
mock_router_instance.cache = DualCache()
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("masked-params-session", request_kwargs),
value={"model": "opus", "tier": "REASONING"},
)
message = "reason carefully about this" if route == "keyword" else "hello"
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": message}],
)
assert response is not None
assert response.litellm_params == params
assert response.routing_decision is not None
assert response.routing_decision["tier_litellm_params"] == {
"reasoning_effort": "xhigh",
"api_key": "secr*******-key",
}
@pytest.mark.asyncio
async def test_session_pin_outside_tiers_does_not_inherit_medium_params(mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": "mini",
"MEDIUM": {"model_name": "medium", "litellm_params": {"reasoning_effort": "low"}},
},
"session_affinity": True,
"default_model": "orphan",
},
)
request_kwargs = {"metadata": {"session_id": "orphan-session"}}
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("orphan-session", request_kwargs),
value="orphan",
)
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.model == "orphan"
assert response.litellm_params == {}
@pytest.mark.asyncio
async def test_session_pin_uses_recorded_tier_when_model_is_in_multiple_tiers(mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": {"model_name": "shared", "litellm_params": {"reasoning_effort": "low"}},
"REASONING": {"model_name": "shared", "litellm_params": {"reasoning_effort": "xhigh"}},
},
"session_affinity": True,
},
)
request_kwargs = {"metadata": {"session_id": "shared-session"}}
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("shared-session", request_kwargs),
value={"model": "shared", "tier": "SIMPLE"},
)
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.litellm_params == {"reasoning_effort": "low"}
assert response.routing_decision is not None
assert response.routing_decision["tier"] == "SIMPLE"
@pytest.mark.asyncio
async def test_session_pin_survives_json_list_round_trip(mock_router_instance):
cache = AsyncMock()
cache.async_get_cache = AsyncMock(return_value=["shared", "SIMPLE"])
mock_router_instance.cache = cache
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": {"model_name": "shared", "litellm_params": {"reasoning_effort": "low"}},
"REASONING": {"model_name": "shared", "litellm_params": {"reasoning_effort": "xhigh"}},
},
"session_affinity": True,
},
)
request_kwargs = {"metadata": {"session_id": "json-round-trip-session"}}
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.model == "shared"
assert response.litellm_params == {"reasoning_effort": "low"}
assert cache.async_set_cache.call_args.kwargs["value"] == {"model": "shared", "tier": "SIMPLE"}

View file

@ -7,6 +7,9 @@ category, it prints `run` or `skip`. The gating contract we lock in here:
* docs-only changes (``*.md``, ``*.mdx``, ``docs/``) run nothing
* client-only changes (``ui/``) run client jobs but skip backend jobs
* any backend change runs both client and backend jobs
* ``ui`` tracks ``ui/`` plus CI config, so a backend-only change skips it
where ``client`` would still run, while a change to the workflows that
define the dashboard jobs still exercises them
If this logic silently regresses, real test jobs get skipped, so these cases
are the guardrail against that.
@ -40,6 +43,7 @@ def classify(category: str, changed: list[str]) -> str:
DOCS = ["README.md", "docs/my_website/index.mdx", "litellm/anywhere.md"]
CLIENT = ["ui/litellm-dashboard/src/App.tsx"]
BACKEND = ["litellm/main.py"]
CI = [".github/workflows/test-litellm-ui-unit.yml"]
@pytest.mark.parametrize(
@ -48,20 +52,27 @@ BACKEND = ["litellm/main.py"]
# docs-only: skip everything
("backend", DOCS, "skip"),
("client", DOCS, "skip"),
("ui", DOCS, "skip"),
("backend", [], "skip"),
("client", [], "skip"),
("ui", [], "skip"),
# client-only: backend skips, client runs
("backend", CLIENT, "skip"),
("client", CLIENT, "run"),
("ui", CLIENT, "run"),
("backend", CLIENT + DOCS, "skip"),
("client", CLIENT + DOCS, "run"),
("ui", CLIENT + DOCS, "run"),
# any backend change: both run ("backend runs both")
("backend", BACKEND, "run"),
("client", BACKEND, "run"),
("ui", BACKEND, "skip"),
("backend", BACKEND + DOCS, "run"),
("client", BACKEND + DOCS, "run"),
("ui", BACKEND + DOCS, "skip"),
("backend", BACKEND + CLIENT, "run"),
("client", BACKEND + CLIENT, "run"),
("ui", BACKEND + CLIENT, "run"),
],
)
def test_classify_decisions(category: str, changed: list[str], expected: str) -> None:
@ -71,6 +82,31 @@ def test_classify_decisions(category: str, changed: list[str], expected: str) ->
def test_markdown_under_ui_counts_as_client_not_docs() -> None:
assert classify("client", ["ui/litellm-dashboard/README.md"]) == "run"
assert classify("backend", ["ui/litellm-dashboard/README.md"]) == "skip"
assert classify("ui", ["ui/litellm-dashboard/README.md"]) == "run"
def test_ci_config_changes_reach_every_category() -> None:
"""A workflow edit has to exercise the jobs it defines, otherwise the change
ships unvalidated: the dashboard jobs would skip on the very pull request
that rewrites them."""
assert classify("ui", CI) == "run"
assert classify("backend", CI) == "run"
assert classify("client", CI) == "run"
def test_markdown_under_dot_github_is_still_docs() -> None:
"""`.github/**` counting as CI config must not drag the pull request template
and other markdown back into running the full suite."""
assert classify("ui", [".github/pull_request_template.md"]) == "skip"
assert classify("backend", [".github/pull_request_template.md"]) == "skip"
def test_ui_and_client_diverge_on_a_backend_only_change() -> None:
"""`client` gates CircleCI's dashboard end-to-end jobs, which drive a real
proxy and so must run on backend changes. `ui` gates the dashboard build and
its unit tests, which cannot see the backend at all."""
assert classify("client", BACKEND) == "run"
assert classify("ui", BACKEND) == "skip"
def test_non_docs_directory_with_docs_in_name_is_backend() -> None:

View file

@ -1,12 +1,13 @@
"""Regression tests for the GitHub Actions change-based job gating.
`.github/scripts/detect_backend_changes.sh` decides whether a pull request's
backend unit-test jobs do real work. It asks the API which files the pull
request touches and hands them to `classify_changes.sh`. The contract locked in
here:
`.github/scripts/detect_changes.sh` decides whether a pull request's jobs do
real work. It asks the API which files the pull request touches and hands them
to `classify_changes.sh` under one category. The contract locked in here:
* a UI-only pull request skips backend jobs even when the checked-out merge
ref carries backend commits from the base branch
* the ui category is the mirror image: it skips when only backend files
changed, so a backend-only PR stops building and unit-testing the dashboard
* anything the classification cannot resolve (no pull request, an API
failure, a truncated file list, a broken classifier) runs the job
"""
@ -19,7 +20,7 @@ import subprocess
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
SCRIPT = REPO_ROOT / ".github" / "scripts" / "detect_backend_changes.sh"
SCRIPT = REPO_ROOT / ".github" / "scripts" / "detect_changes.sh"
CLASSIFIER = REPO_ROOT / ".circleci" / "scripts" / "classify_changes.sh"
UI_FILE = "ui/litellm-dashboard/src/components/Teams.tsx"
@ -84,6 +85,7 @@ def _run(
changed_file_count: str | None = None,
gh_exit_code: int = 0,
classifier_body: str | None = None,
category: str | None = None,
) -> tuple[str, str]:
"""Run the script against a stubbed `gh`; returns (decision, stdout)."""
bin_dir = tmp_path / "bin"
@ -102,6 +104,10 @@ def _run(
env["REPO"] = "BerriAI/litellm"
env["PR_NUMBER"] = pr_number
env["CHANGED_FILE_COUNT"] = changed_file_count if changed_file_count is not None else str(len(files))
if category is not None:
env["CATEGORY"] = category
else:
env.pop("CATEGORY", None)
tree = _scripts_tree(tmp_path, classifier_body)
result = subprocess.run(
@ -187,3 +193,43 @@ def test_unexpected_classifier_output_runs(tmp_path: Path) -> None:
)
assert decision == "decision=run"
assert "unexpected decision: maybe" in stdout
def test_ui_category_skips_a_backend_only_pr(tmp_path: Path) -> None:
"""The dashboard build and its unit tests cannot be affected by a pull request
that touches no `ui/` file, and the `client` category cannot express that
because it deliberately runs whenever the backend changes."""
decision, _ = _run(tmp_path, files=[BACKEND_FILE], category="ui")
assert decision == "decision=skip"
def test_ui_category_runs_a_ui_only_pr(tmp_path: Path) -> None:
decision, _ = _run(tmp_path, files=[UI_FILE], category="ui")
assert decision == "decision=run"
def test_ui_category_runs_a_mixed_pr(tmp_path: Path) -> None:
decision, _ = _run(tmp_path, files=[UI_FILE, BACKEND_FILE], category="ui")
assert decision == "decision=run"
def test_absent_category_still_runs_a_backend_pr(tmp_path: Path) -> None:
"""Callers that pass no category keep the pre-existing backend behaviour."""
assert _run(tmp_path, files=[BACKEND_FILE])[0] == "decision=run"
def test_absent_category_still_skips_a_ui_pr(tmp_path: Path) -> None:
assert _run(tmp_path, files=[UI_FILE])[0] == "decision=skip"
def test_ui_category_fails_open_when_the_api_fails(tmp_path: Path) -> None:
decision, stdout = _run(tmp_path, files=[], gh_exit_code=1, category="ui")
assert decision == "decision=run"
assert "detect-changes[ui]" in stdout
def test_ui_category_runs_when_the_ui_workflows_themselves_change(tmp_path: Path) -> None:
"""Without this the dashboard jobs would skip on the pull request that edits
them, shipping a workflow change nothing ever exercised."""
decision, _ = _run(tmp_path, files=[".github/workflows/test-litellm-ui-unit.yml"], category="ui")
assert decision == "decision=run"

View file

@ -1584,3 +1584,116 @@ def test_inherit_builtin_tiered_output_rate_leaves_a_user_rate_alone():
)
assert model_info["output_cost_per_token"] == 9e-07
# --- a config.yaml PTU deployment must not also bill per token ------------------
_PTU_MODEL_INFO = {
"team_id": "team-alpha",
"ptu_count": 100,
"cost_per_ptu_per_hour": 0.02,
"ptu_effective_from": "2026-01-01T00:00:00Z",
}
def _ptu_router(model_info=None, litellm_params=None, ptu_enabled=True):
"""A router built the way loading config.yaml builds one."""
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True" if ptu_enabled else ""}, clear=False):
return Router(
model_list=[
{
"model_name": "gpt-4o-ptu",
"litellm_params": {
"model": "anthropic/claude-sonnet-4-5-20250929",
"api_key": "sk-not-used",
**(litellm_params or {}),
},
"model_info": dict(_PTU_MODEL_INFO if model_info is None else model_info),
}
]
)
def test_a_config_ptu_deployment_bills_nothing_per_token():
"""Reserved capacity is already billed by the hour, so charging its traffic bills the
same tokens twice. Left unset the rate falls back to the public cost map, which makes
the double charge the default rather than an opt-in."""
router = _ptu_router(litellm_params={"input_cost_per_token": 5e-06, "output_cost_per_token": 1.5e-05})
entry = router.model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 0.0
assert entry["litellm_params"]["output_cost_per_token"] == 0.0
assert entry["model_info"]["input_cost_per_token"] == 0.0
assert litellm.model_cost[entry["model_info"]["id"]]["input_cost_per_token"] == 0.0
@pytest.mark.parametrize(
"backend",
["anthropic/claude-sonnet-4-5-20250929", "azure/gpt-4o", "gemini/gemini-2.5-flash"],
)
def test_a_config_ptu_deployment_imports_no_cache_rate_from_its_backend(backend):
"""The cache back-fill runs whenever input_cost_per_token is set, and 0.0 is set, so a
partially zeroed deployment would silently inherit the backend model's real cache rates.
Every backend here publishes non-zero ones, which is what makes the assertion mean
something."""
cache_fields = (
"cache_creation_input_token_cost",
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost",
"cache_read_input_token_cost_above_200k_tokens",
)
builtin = litellm.get_model_info(model=backend)
assert any(builtin.get(field) for field in cache_fields), "backend publishes no cache pricing to leak"
router = _ptu_router(litellm_params={"model": backend})
priced = litellm.model_cost[router.model_list[0]["model_info"]["id"]]
assert [field for field in cache_fields if priced.get(field)] == []
def test_zeroing_a_ptu_deployment_leaves_its_backend_model_priced():
"""A sibling deployment on the same backend must keep billing normally."""
backend = "anthropic/claude-sonnet-4-5-20250929"
builtin = litellm.get_model_info(model=backend)["input_cost_per_token"]
assert builtin > 0
_ptu_router(litellm_params={"model": backend})
assert litellm.get_model_info(model=backend)["input_cost_per_token"] == builtin
def test_zeroing_does_not_change_the_deployment_id():
"""The id is a hash of the deployment's params and keys its cooldowns, its budget, and
every spend row already written against it."""
params = {"input_cost_per_token": 5e-06}
priced = _ptu_router(litellm_params=params, ptu_enabled=False).model_list[0]["model_info"]["id"]
zeroed = _ptu_router(litellm_params=params).model_list[0]["model_info"]["id"]
assert priced == zeroed
def test_a_database_backed_deployment_is_left_alone():
"""The write endpoints already zero those, and they answer 400 rather than silently
rewriting a rate the caller sent."""
entry = _ptu_router(model_info={**_PTU_MODEL_INFO, "db_model": True}).model_list[0]
assert entry["litellm_params"].get("input_cost_per_token") is None
def test_nothing_is_zeroed_while_the_feature_is_off():
"""No flat cost accrues with the flag off, so zeroing would serve the traffic free."""
entry = _ptu_router(litellm_params={"input_cost_per_token": 5e-06}, ptu_enabled=False).model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 5e-06
@pytest.mark.parametrize("dropped", ["team_id", "ptu_effective_from"], ids=["no team_id", "no ptu_effective_from"])
def test_a_deployment_the_rollup_will_not_charge_is_not_zeroed(dropped):
"""The rollup refuses to price a reservation missing either field, so zeroing on the
looser count-and-rate test alone would leave the deployment serving for free with
nothing charged in its place."""
incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped}
entry = _ptu_router(model_info=incomplete, litellm_params={"input_cost_per_token": 5e-06}).model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 5e-06

View file

@ -12,13 +12,13 @@ Most of the suite predates this split and is not yet classified, so an unsuffixe
Assert something the user could perceive, and assert it precisely enough that the test fails when the behaviour breaks. `eslint-plugin-testing-library` and `eslint-plugin-jest-dom` enforce the mechanical part of that. Two of the enabled rules exist because the failure they catch is silent rather than cosmetic: `await-async-queries` catches an unawaited `findBy*`, whose returned Promise is always truthy and makes the whole assertion vacuous, and `no-wait-for-side-effects` catches work inside a `waitFor` callback, which is retried on every poll. Prefer `findBy*` over `waitFor` wrapped around `getBy*`, and keep a `waitFor` callback to a single assertion
Do not trust `eslint --fix` for these two plugins. Fixing the suite in bulk produced seven distinct kinds of broken output. Four fail loudly: `no-wait-for-side-effects` and `no-wait-for-multiple-assertions` hoist a statement out of the `waitFor` callback while leaving the `const` it reads inside, `prefer-enabled-disabled` drops a closing paren when the subject carries a type assertion, `prefer-presence-queries` swaps in a query it never destructures, and `prefer-in-document` collapses `getAllBy*` to `getBy*` on a value still indexed as an array. Two fail quietly, which is worse: `prefer-checked` swaps the `checked` attribute for the `.checked` property, and antd radios set one without the other, and `prefer-to-have-text-content` wraps arbitrary strings in `new RegExp()` without escaping, so `toContain("100K+ requests")` becomes a pattern meaning "100 followed by one-or-more K". That last one compiles, lints clean, and keeps passing while no longer asserting what it says. Pass a plain string to `toHaveTextContent`, which is already a substring match. Run the fixer on a handful of files at a time and read the diff
Do not trust `eslint --fix` for these two plugins. Fixing the suite in bulk produced seven distinct kinds of broken output. Four fail loudly: `no-wait-for-side-effects` and `no-wait-for-multiple-assertions` hoist a statement out of the `waitFor` callback while leaving the `const` it reads inside, `prefer-enabled-disabled` drops a closing paren when the subject carries a type assertion, `prefer-presence-queries` swaps in a query it never destructures, and `prefer-in-document` collapses `getAllBy*` to `getBy*` on a value still indexed as an array. Two fail quietly, which is worse: `prefer-checked` swaps the `checked` attribute for the `.checked` property, and a radio can set one without the other, and `prefer-to-have-text-content` wraps arbitrary strings in `new RegExp()` without escaping, so `toContain("100K+ requests")` becomes a pattern meaning "100 followed by one-or-more K". That last one compiles, lints clean, and keeps passing while no longer asserting what it says. Pass a plain string to `toHaveTextContent`, which is already a substring match. Run the fixer on a handful of files at a time and read the diff
`jest-dom/prefer-to-have-value` stays off because its fixer is wrong here, not merely noisy. It matches any attribute whose name contains "value", so it rewrites `toHaveAttribute("aria-valuenow", n)` into `toHaveValue(n)`, and jest-dom's `toHaveValue` only supports form controls, so the assertion fails on the `role="meter"` elements the dashboard renders. Assert ARIA value attributes with `toHaveAttribute`
Reach for `fireEvent.change` rather than `user.type` when a test only needs a field to hold a value. `user.type` dispatches one event per character and re-renders the whole form each time, which is why a single form test could burn seven seconds. Keep `user.type` where the typing itself is the behaviour under test: an autocomplete that filters per keystroke, a debounce, a key handler, or any Base UI combobox, whose filter state is driven by real keyboard input and does not react to a raw change event
A test may reach for a component library's own CSS class only when that library exposes no role, label, title or ARIA state to query instead, and then the line carries a suppression naming the rule and the reason. Check first: antd icons render as `role="img"` with an `aria-label`, and antd `Form.Item` associates its label with the control, so both are reachable accessibly. When a label does not resolve, suspect the control rather than the test, since a custom wrapper that destructures props without spreading them drops the `id` antd injects and leaves the rendered label pointing at nothing
A test may reach for a component library's own CSS class only when that library exposes no role, label, title or ARIA state to query instead. Check first: the shadcn primitives forward roles and `aria-label`, and the shared form field associates its label with the control, so both are reachable accessibly. When nothing accessible identifies the element, prefer its `data-slot` attribute, which the primitives set deliberately and treat as stable. When a label does not resolve, suspect the control rather than the test, since a custom wrapper that destructures props without spreading them drops the `id` the field generates and leaves the rendered label pointing at nothing
Rules beyond the enabled set were measured against the whole suite and left off rather than recorded in a budget file, because a ceiling that permits a violation anywhere is worse than an honest gap. `no-node-access` and `no-container` are the ones worth revisiting first, since they catch the DOM archaeology the rules above only discourage. `prefer-implicit-assert` and `prefer-explicit-assert` contradict each other, so neither is enabled

View file

@ -5,9 +5,6 @@
}
},
"src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx": {
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -22,9 +19,6 @@
"no-nested-ternary": {
"count": 3
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 2
}
@ -57,9 +51,6 @@
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 1
}
@ -199,11 +190,6 @@
"count": 1
}
},
"src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx": {
"no-nested-ternary": {
"count": 3
@ -543,18 +529,10 @@
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": {
"no-nested-ternary": {
"count": 1
@ -570,11 +548,6 @@
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -782,11 +755,6 @@
"count": 1
}
},
"src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": {
"local/no-complex-jsx-arrow": {
"count": 2
@ -894,9 +862,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 2
},
@ -987,9 +952,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 1
}
@ -1023,9 +985,6 @@
"src/app/(dashboard)/prompts/_components/add_prompt_form.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/prompts/_components/index.tsx": {
@ -1213,9 +1172,6 @@
}
},
"src/app/(dashboard)/users/_components/BulkEditUsers.tsx": {
"no-restricted-imports": {
"count": 1
},
"prefer-const": {
"count": 1
}
@ -1225,11 +1181,6 @@
"count": 1
}
},
"src/app/(dashboard)/users/_components/user_edit_view.test.tsx": {
"no-nested-ternary": {
"count": 1
}
},
"src/app/(dashboard)/users/_components/user_edit_view.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1257,11 +1208,6 @@
"count": 1
}
},
"src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx": {
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx": {
"no-nested-ternary": {
"count": 2
@ -1354,9 +1300,6 @@
}
},
"src/components/CreateUserButton.tsx": {
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1400,18 +1343,10 @@
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
},
"src/components/Settings/AdminSettings/PluginSettings/PluginSettings.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx": {
"max-nested-callbacks": {
"count": 1
@ -1437,18 +1372,7 @@
"count": 1
}
},
"src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": {
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1473,9 +1397,6 @@
"no-nested-ternary": {
"count": 2
},
"no-restricted-imports": {
"count": 1
},
"prefer-const": {
"count": 2
},
@ -1488,11 +1409,6 @@
"count": 1
}
},
"src/components/UsagePage/components/EntityUsage/TopKeyView.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/UsagePage/utils/value_formatters.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1511,20 +1427,9 @@
"count": 1
}
},
"src/components/add_model/AddModelForm.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/AddModelForm.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/RouterConfigBuilder.tsx": {
@ -1535,9 +1440,6 @@
"src/components/add_model/add_auto_router_tab.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/add_model_modes.tsx": {
@ -1548,9 +1450,6 @@
"src/components/add_model/advanced_settings.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/auto_router_connection_test.tsx": {
@ -1570,9 +1469,6 @@
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 2
}
@ -1593,9 +1489,6 @@
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/add_model/model_connection_test.tsx": {
@ -1613,9 +1506,6 @@
"no-nested-ternary": {
"count": 3
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 3
}
@ -1625,15 +1515,7 @@
"count": 1
}
},
"src/components/agent_management/AgentSelector.test.tsx": {
"react/display-name": {
"count": 1
}
},
"src/components/agent_management/AgentSelector.tsx": {
"no-restricted-imports": {
"count": 1
},
"prefer-const": {
"count": 1
}
@ -1714,9 +1596,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"no-restricted-syntax": {
"count": 3
},
@ -1724,11 +1603,6 @@
"count": 1
}
},
"src/components/common_components/AccessGroupSelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/DeleteResourceModal.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -1739,16 +1613,6 @@
"count": 1
}
},
"src/components/common_components/MetadataKeyValueFields.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/MetadataKeyValueFields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/ModelAliasManager.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -1759,11 +1623,6 @@
"count": 1
}
},
"src/components/common_components/RateLimitTypeFormItem.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/budget_duration_dropdown.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1772,9 +1631,6 @@
"src/components/common_components/check_openapi_schema.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/common_components/fetch_teams.tsx": {
@ -1845,16 +1701,6 @@
"count": 1
}
},
"src/components/key_team_helpers/BudgetFallbacksEditor.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/key_team_helpers/BudgetWindowsEditor.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/key_team_helpers/fetch_available_models_team_key.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1923,11 +1769,6 @@
"count": 1
}
},
"src/components/mcp_tools/ByokCredentialModal.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/mcp_tools/MCPToolArgumentsForm.tsx": {
"no-nested-ternary": {
"count": 1
@ -1943,21 +1784,11 @@
"count": 1
}
},
"src/components/model_add/CredentialModal.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/model_add/reuse_credentials.tsx": {
"local/filename-pascal-case": {
"count": 1
}
},
"src/components/model_dashboard/ModelSettingsModal/ModelSettingsModal.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/model_filters.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2020,9 +1851,6 @@
"src/components/onboarding_link.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/organisms/create_key_button.tsx": {
@ -2032,9 +1860,6 @@
"max-lines": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"prefer-const": {
"count": 2
},
@ -2111,11 +1936,6 @@
"count": 1
}
},
"src/components/router_settings/RoutingStrategySelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/router_settings/index.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2128,9 +1948,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/preserve-manual-memoization": {
"count": 1
}
@ -2264,11 +2081,6 @@
"count": 1
}
},
"src/components/team/LoggingSettings.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/team/TeamInfo.tsx": {
"max-lines": {
"count": 1
@ -2276,9 +2088,6 @@
"no-nested-ternary": {
"count": 3
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -2503,9 +2312,6 @@
"src/components/update_model_credentials_modal.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/user_agent_activity.tsx": {
@ -2614,11 +2420,6 @@
"count": 2
}
},
"src/contexts/AntdGlobalProvider.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/contexts/AuthContext.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1

View file

@ -57,15 +57,6 @@ const eslintConfig = [
message:
"@tremor/react is being phased out; build new UI with shadcn/ui primitives instead of adding tremor imports.",
},
{
group: ["antd", "antd/*"],
message:
"antd is being phased out; build new UI with shadcn/ui primitives instead of adding antd imports.",
},
{
group: ["@ant-design/icons", "@ant-design/icons/*"],
message: "@ant-design/icons is gone from the dashboard; use lucide-react instead.",
},
],
},
],
@ -94,7 +85,6 @@ const eslintConfig = [
files: ["src/**/*.test.{ts,tsx}", "tests/**/*.{ts,tsx}"],
plugins: { "testing-library": testingLibrary, "jest-dom": jestDom },
rules: {
"local/no-antd-class-selectors": "error",
"testing-library/await-async-queries": "error",
"testing-library/no-wait-for-multiple-assertions": "error",
"testing-library/no-wait-for-side-effects": "error",

File diff suppressed because it is too large Load diff

View file

@ -24,7 +24,6 @@
"gen:api": "node scripts/gen-api-types.mjs"
},
"dependencies": {
"@ant-design/cssinjs": "1.24.0",
"@anthropic-ai/sdk": "0.92.0",
"@base-ui/react": "^1.6.0",
"@headlessui/tailwindcss": "0.2.2",
@ -34,7 +33,6 @@
"@tanstack/react-query": "5.100.7",
"@tanstack/react-table": "8.21.3",
"@types/papaparse": "5.5.2",
"antd": "5.29.3",
"cva": "1.0.0-beta.4",
"date-fns": "^4.4.0",
"dayjs": "1.11.19",

View file

@ -2,7 +2,6 @@ import noLargeInlineObjectArg from "./no-large-inline-object-arg.mjs";
import noLongConditionChain from "./no-long-condition-chain.mjs";
import noComplexJsxArrow from "./no-complex-jsx-arrow.mjs";
import filenamePascalCase from "./filename-pascal-case.mjs";
import noAntdClassSelectors from "./no-antd-class-selectors.mjs";
const plugin = {
rules: {
@ -10,7 +9,6 @@ const plugin = {
"no-long-condition-chain": noLongConditionChain,
"no-complex-jsx-arrow": noComplexJsxArrow,
"filename-pascal-case": filenamePascalCase,
"no-antd-class-selectors": noAntdClassSelectors,
},
};

View file

@ -1,45 +0,0 @@
const SELECTOR_REFERENCE = /\.(?:ant|anticon)-[a-z0-9-]+/;
const BARE_CLASS_REFERENCE = /^(?:ant|anticon)-[a-z0-9-]+$/;
const CLASS_ASSERTION_CALLEES = new Set(["toHaveClass", "contains", "toContain"]);
const isClassAssertionArgument = (node) => {
const call = node.parent;
if (call?.type !== "CallExpression" || !call.arguments.includes(node)) return false;
const callee = call.callee;
return callee?.type === "MemberExpression" && CLASS_ASSERTION_CALLEES.has(callee.property?.name);
};
const rule = {
meta: {
type: "problem",
docs: {
description:
"Disallow locating or asserting on antd's internal CSS classes in tests; query by role, label or text instead.",
},
schema: [],
messages: {
antdClass:
'Test depends on antd internal class "{{value}}". Query by role, label or text (getByLabelText, getByRole("combobox"), getByTitle) so the test survives the shadcn migration.',
},
},
create(context) {
const report = (node, value) => {
if (typeof value !== "string") return;
const matches =
SELECTOR_REFERENCE.test(value) || (BARE_CLASS_REFERENCE.test(value) && isClassAssertionArgument(node));
if (!matches) return;
context.report({ node, messageId: "antdClass", data: { value } });
};
return {
Literal(node) {
report(node, node.value);
},
TemplateElement(node) {
report(node, node.value.cooked);
},
};
},
};
export default rule;

View file

@ -3,7 +3,7 @@ import { useDeleteAccessGroup } from "@/app/(dashboard)/hooks/accessGroups/useDe
import { Plus, SearchIcon, X } from "lucide-react";
import { useMemo, useState } from "react";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import { PageHeader } from "@/components/shared/PageHeader";
import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader";
import { Button } from "@/components/ui/button";
import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group";
import { AccessGroupDetail } from "./AccessGroupsDetailsPage";
@ -61,7 +61,7 @@ export function AccessGroupsPage() {
return (
<div className="p-6 px-12">
<div className="mb-4">
<PageHeader
<LegacyPageHeader
title="Access Groups"
subtitle="Manage resource permissions for your organization"
actions={

View file

@ -22,8 +22,8 @@ interface ResourceTone {
}
const RESOURCE_TONES: Record<"models" | "mcpServers" | "agents", ResourceTone> = {
models: { icon: Layers, className: "bg-blue-50 text-blue-700 ring-blue-600/20" },
mcpServers: { icon: Server, className: "bg-cyan-50 text-cyan-700 ring-cyan-600/20" },
models: { icon: Layers, className: "bg-info/10 text-info ring-blue-600/20" },
mcpServers: { icon: Server, className: "bg-info/10 text-info ring-cyan-600/20" },
agents: { icon: Bot, className: "bg-purple-50 text-purple-700 ring-purple-600/20" },
};

View file

@ -7,7 +7,7 @@ import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { Button } from "@/components/ui/button";
import { Card } from "@/components/ui/card";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { Space, Tabs, Typography } from "antd";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { Info, TriangleAlert } from "lucide-react";
import React, { useEffect, useState } from "react";
import NewBadge from "@/components/common_components/NewBadge";
@ -35,8 +35,6 @@ import { Input } from "@/components/ui/input";
import { useZodForm } from "@/lib/forms/useZodForm";
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
const { Title, Paragraph, Text } = Typography;
const allowedIPSchema = z.object({
ip: z.string().min(1, "Please enter an IP address"),
});
@ -223,7 +221,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
children: (
<>
<Card className="block p-6">
<Title level={4}> ✨ Security Settings</Title>
<h3 className="mb-2 text-base font-semibold text-foreground">✨ Security Settings</h3>
<Alert variant="warning">
<TriangleAlert />
<AlertTitle>SSO Configuration Deprecated</AlertTitle>
@ -329,7 +327,9 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
<DialogHeader>
<DialogTitle>Confirm Delete</DialogTitle>
</DialogHeader>
<Text>Are you sure you want to delete the IP address: {ipToDelete}?</Text>
<span className="text-sm text-foreground">
Are you sure you want to delete the IP address: {ipToDelete}?
</span>
<DialogFooter>
<Button className="mx-1" onClick={() => confirmDeleteIP()}>
Yes
@ -379,11 +379,10 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
{
key: "ui-settings",
label: (
<Space>
<Text>
UI Settings <NewBadge />
</Text>
</Space>
<span className="flex items-center gap-1.5">
UI Settings
<NewBadge />
</span>
),
children: (
<div className="flex flex-col gap-4">
@ -411,9 +410,22 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
return (
<div className="w-full m-2 mt-2 p-8">
<Title level={4}>Admin Access </Title>
<Paragraph>Go to &apos;Internal Users&apos; page to add other admins.</Paragraph>
<Tabs items={tabItems} />
<h2 className="mb-2 text-base font-semibold text-foreground">Admin Access</h2>
<p className="mb-4 text-sm text-foreground">Go to &apos;Internal Users&apos; page to add other admins.</p>
<Tabs defaultValue={tabItems[0].key}>
<TabsList variant="line" className="mb-4 h-auto flex-wrap">
{tabItems.map((item) => (
<TabsTrigger key={item.key} value={item.key} className="flex-none">
{item.label}
</TabsTrigger>
))}
</TabsList>
{tabItems.map((item) => (
<TabsContent key={item.key} value={item.key}>
{item.children}
</TabsContent>
))}
</Tabs>
</div>
);
};

View file

@ -73,7 +73,7 @@ const AgentsTable: React.FC<AgentsTableProps> = ({
render={
<div className="flex items-center gap-2">
<CircleCheck
className={healthCheckEnabled ? "size-4 text-green-500" : "size-4 text-muted-foreground"}
className={healthCheckEnabled ? "size-4 text-success" : "size-4 text-muted-foreground"}
/>
<span className="text-sm text-muted-foreground">Health Check</span>
<Switch

View file

@ -1,5 +1,6 @@
import React from "react";
import { render, screen, fireEvent, within } from "@testing-library/react";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import AddAgentForm from "./add_agent_form";
import * as networking from "@/components/networking";
@ -50,7 +51,7 @@ describe("AddAgentForm logos", () => {
expect(titleLogo).toBeInstanceOf(HTMLImageElement);
expect(titleLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
const selectionLogo = await screen.findByAltText("A2A Agent logo");
const selectionLogo = within(await screen.findByRole("combobox")).getByAltText("A2A Agent logo");
expect(selectionLogo).toBeInstanceOf(HTMLImageElement);
expect(selectionLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
});
@ -64,10 +65,12 @@ describe("AddAgentForm logos", () => {
});
it("renders the option logo when the agent type dropdown is opened", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderForm();
await screen.findByAltText("A2A Agent logo");
fireEvent.mouseDown(screen.getByRole("combobox"));
const trigger = await screen.findByRole("combobox");
await within(trigger).findByAltText("A2A Agent logo");
await user.click(trigger);
const optionLogos = await screen.findAllByAltText("A2A Agent logo");
expect(optionLogos.length).toBeGreaterThanOrEqual(2);
@ -88,9 +91,9 @@ describe("AddAgentForm logos", () => {
expect(screen.queryByAltText("Agent logo")).not.toBeInTheDocument();
expect(within(header).getByText("A")).toBeInTheDocument();
const selectionLogo = screen.getByAltText("A2A Agent logo");
fireEvent.error(selectionLogo);
expect(screen.queryByAltText("A2A Agent logo")).not.toBeInTheDocument();
const trigger = screen.getByRole("combobox");
fireEvent.error(within(trigger).getByAltText("A2A Agent logo"));
expect(within(trigger).queryByAltText("A2A Agent logo")).not.toBeInTheDocument();
expect(warnSpy).toHaveBeenCalledTimes(2);
warnSpy.mockRestore();
});

View file

@ -1,14 +1,15 @@
import React, { useState, useEffect } from "react";
import { Select, Steps, Tag } from "antd";
import { FormProvider, useForm, useWatch } from "react-hook-form";
import { toast } from "@/lib/toast";
import { Logo } from "@/components/molecules/logo/Logo";
import { Bot, CircleCheck, Key, LayoutGrid } from "lucide-react";
import { Bot, Check, CircleCheck, Key, LayoutGrid } from "lucide-react";
import CreatedKeyDisplay from "@/components/shared/CreatedKeyDisplay";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { Input } from "@/components/ui/input";
import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group";
import { Select, SelectContent, SelectItem, SelectSeparator, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Separator } from "@/components/ui/separator";
import { Switch } from "@/components/ui/switch";
import { Textarea } from "@/components/ui/textarea";
@ -51,10 +52,61 @@ import MCPToolPermissions from "@/components/mcp_server_management/MCPToolPermis
import GuardrailSelector from "@/components/guardrails/GuardrailSelector";
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
const { Step } = Steps;
const CUSTOM_AGENT_TYPE = "custom";
const STEP_TITLES = ["Configure", "Entitlements", "Governance", "Agent Management", "Ready"] as const;
const stepMarkerClass = (index: number, current: number): string => {
if (index < current) return "border-primary text-primary";
if (index === current) return "border-primary bg-primary text-primary-foreground";
return "border-border text-muted-foreground";
};
const stepTitleClass = (index: number, current: number): string => {
if (index === current) return "font-medium text-foreground";
if (index < current) return "text-foreground";
return "text-muted-foreground";
};
const AgentTypeLabel: React.FC<{ agentType: string; info: AgentCreateInfo | undefined }> = ({ agentType, info }) => {
if (agentType === CUSTOM_AGENT_TYPE) {
return (
<span className="flex items-center gap-2">
<LayoutGrid className="size-4 text-warning" />
<span>Custom / Other</span>
</span>
);
}
if (!info) return <>{agentType}</>;
return (
<span className="flex items-center gap-2">
<Logo src={info.logo_url} label={info.agent_type_display_name} className="h-4 w-4 object-contain" />
<span>{info.agent_type_display_name}</span>
</span>
);
};
const StepProgress: React.FC<{ current: number }> = ({ current }) => (
<ol aria-label="Agent creation steps" className="mb-8 flex items-center">
{STEP_TITLES.map((title, index) => (
<li
key={title}
aria-current={index === current ? "step" : undefined}
className="flex flex-1 items-center gap-2 last:flex-none"
>
<span
aria-hidden="true"
className={`flex size-6 shrink-0 items-center justify-center rounded-full border text-xs ${stepMarkerClass(index, current)}`}
>
{index < current ? <Check className="size-3.5" /> : index + 1}
</span>
<span className={`text-xs whitespace-nowrap ${stepTitleClass(index, current)}`}>{title}</span>
{index < STEP_TITLES.length - 1 && <span aria-hidden="true" className="mx-2 h-px flex-1 bg-border" />}
</li>
))}
</ol>
);
const SHARED_INITIAL_VALUES: AgentFormValues = {
allowed_mcp_servers_and_groups: { servers: [], accessGroups: [] },
mcp_tool_permissions: {},
@ -536,7 +588,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<h4 className="mb-3 text-sm font-medium text-foreground">Budgets &amp; Rate Limits</h4>
<div className="space-y-4">
{!requireTraceIdOutbound && (
<div className="rounded-lg border border-yellow-200 bg-yellow-50 p-3 text-sm text-yellow-800 dark:border-yellow-800 dark:bg-yellow-950 dark:text-yellow-300">
<div className="rounded-lg border border-warning/20 bg-warning/10 p-3 text-sm text-warning">
Enable &quot;Require x-litellm-trace-id on calls BY this agent&quot; in Tracing to configure budgets and
rate limits.
</div>
@ -684,66 +736,45 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<FieldLabel htmlFor="agent-type">
{labelWithHint("Agent Type", "Select the type of agent you want to create")}
</FieldLabel>
<Select
id="agent-type"
value={agentType}
onChange={handleAgentTypeChange}
size="large"
style={{ width: "100%" }}
optionLabelProp="label"
dropdownRender={(menu) => (
<>
{menu}
<Separator className="my-1" />
<div className="px-2 py-1">
<div className="mb-1 px-2 text-xs font-medium tracking-wide text-muted-foreground uppercase">
Not listed?
</div>
<div
className={`flex cursor-pointer items-center gap-3 rounded px-2 py-2 transition-colors ${
agentType === CUSTOM_AGENT_TYPE
? "bg-amber-50 dark:bg-amber-950"
: "hover:bg-amber-50 dark:hover:bg-amber-950"
}`}
onClick={() => handleAgentTypeChange(CUSTOM_AGENT_TYPE)}
>
<LayoutGrid className="size-4.5 text-amber-600 dark:text-amber-400" />
<div className="flex-1">
<div className="flex items-center gap-2">
<span className="font-medium text-amber-700 dark:text-amber-400">Custom / Other</span>
<Tag color="orange" style={{ fontSize: 10, padding: "0 4px" }}>
GENERIC
</Tag>
</div>
<div className="text-xs text-amber-600 dark:text-amber-400">
For agents that don&apos;t follow a standard protocol, just needs a virtual key
</div>
</div>
</div>
</div>
</>
)}
>
{agentTypeMetadata.map((info) => (
<Select.Option
key={info.agent_type}
value={info.agent_type}
label={
<div className="flex items-center gap-2">
<Logo src={info.logo_url} label={info.agent_type_display_name} className="h-4 w-4 object-contain" />
<span>{info.agent_type_display_name}</span>
</div>
}
>
<div className="flex items-center gap-3 py-1">
<Logo src={info.logo_url} label={info.agent_type_display_name} className="h-5 w-5 object-contain" />
<div>
<div className="font-medium">{info.agent_type_display_name}</div>
{info.description && <div className="text-xs text-muted-foreground">{info.description}</div>}
</div>
</div>
</Select.Option>
))}
<Select value={agentType} onValueChange={(value) => value !== null && handleAgentTypeChange(value)}>
<SelectTrigger id="agent-type" className="h-10 w-full">
<SelectValue>{() => <AgentTypeLabel agentType={agentType} info={selectedAgentTypeInfo} />}</SelectValue>
</SelectTrigger>
<SelectContent className="p-1">
{agentTypeMetadata.map((info) => (
<SelectItem key={info.agent_type} value={info.agent_type}>
<span className="flex items-center gap-3 py-1">
<Logo src={info.logo_url} label={info.agent_type_display_name} className="h-5 w-5 object-contain" />
<span className="block">
<span className="block font-medium">{info.agent_type_display_name}</span>
{info.description && (
<span className="block text-xs text-muted-foreground">{info.description}</span>
)}
</span>
</span>
</SelectItem>
))}
<SelectSeparator />
<div className="mb-1 px-2 text-xs font-medium tracking-wide text-muted-foreground uppercase">
Not listed?
</div>
<SelectItem value={CUSTOM_AGENT_TYPE} className="focus:bg-warning/10 dark:focus:**:text-amber-400">
<span className="flex items-center gap-3">
<LayoutGrid className="size-4.5 shrink-0 text-warning" />
<span className="block">
<span className="flex items-center gap-2">
<span className="font-medium text-warning">Custom / Other</span>
<Badge variant="warning" className="h-4 px-1 text-[10px]">
GENERIC
</Badge>
</span>
<span className="block text-xs whitespace-normal text-warning">
For agents that don&apos;t follow a standard protocol, just needs a virtual key
</span>
</span>
</span>
</SelectItem>
</SelectContent>
</Select>
</Field>
@ -846,10 +877,10 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<div>
{/* Agent name chip */}
<div className="mb-6 flex justify-center">
<Tag color="purple" className="inline-flex items-center gap-1.5 px-3 py-1 text-sm">
<Badge className="h-auto gap-1.5 bg-purple-100 px-3 py-1 text-sm text-purple-700">
<Bot className="size-3.5" />
{agentName}
</Tag>
</Badge>
</div>
<AgentFormField
@ -875,7 +906,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<div
className={`cursor-pointer rounded-lg border-2 p-4 transition-colors ${
keyAssignOption === "create_new"
? "border-indigo-600 bg-indigo-50 dark:border-indigo-400 dark:bg-indigo-950"
? "border-indigo-600 bg-indigo-50"
: "border-border bg-background hover:border-muted-foreground/40"
}`}
onClick={() => setKeyAssignOption("create_new")}
@ -885,7 +916,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<RadioGroupItem value="create_new" aria-label="Create a new key for this agent" />
<div className="flex-1">
<div className="flex items-center gap-2">
<Key className="size-4 text-indigo-600 dark:text-indigo-400" />
<Key className="size-4 text-indigo-600" />
<span className="font-medium text-foreground">Create a new key for this agent</span>
</div>
<p className="mt-1 text-sm text-muted-foreground">A dedicated key scoped to this agent.</p>
@ -904,7 +935,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
)}
</div>
</div>
<Tag color="green">Recommended</Tag>
<Badge variant="success">Recommended</Badge>
</div>
</div>
@ -912,7 +943,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<div
className={`cursor-pointer rounded-lg border-2 p-4 transition-colors ${
keyAssignOption === "existing_key"
? "border-indigo-600 bg-indigo-50 dark:border-indigo-400 dark:bg-indigo-950"
? "border-indigo-600 bg-indigo-50"
: "border-border bg-background hover:border-muted-foreground/40"
}`}
onClick={() => setKeyAssignOption("existing_key")}
@ -959,13 +990,13 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
const renderReadyStep = () => (
<div className="py-6 text-center">
<CircleCheck className="mb-4 size-12 text-green-500" />
<CircleCheck className="mb-4 size-12 text-success" />
<h3 className="mb-2 text-xl font-semibold text-foreground">Agent Created!</h3>
<div className="mb-4 flex justify-center">
<Tag color="purple" className="inline-flex items-center gap-1.5 px-3 py-1 text-sm">
<Badge className="h-auto gap-1.5 bg-purple-100 px-3 py-1 text-sm text-purple-700">
<Bot className="size-3.5" />
{createdAgentName}
</Tag>
</Badge>
</div>
{createdKeyValue && (
<div className="mx-auto mt-4 max-w-md text-left">
@ -998,13 +1029,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
</DialogHeader>
<TooltipProvider>
<div className="mt-4">
<Steps current={currentStep} size="small" className="mb-8">
<Step title="Configure" />
<Step title="Entitlements" />
<Step title="Governance" />
<Step title="Agent Management" />
<Step title="Ready" />
</Steps>
<StepProgress current={currentStep} />
<FormProvider {...form}>
<form onSubmit={(event) => event.preventDefault()} className="space-y-4">

View file

@ -334,7 +334,7 @@ const AgentCardDiscovery: React.FC<AgentCardDiscoveryProps> = ({
{card && (
<div className="mt-4 rounded-lg border border-border bg-background p-4">
<div className="mb-3 flex flex-wrap items-center gap-2">
<CircleCheck className="size-4 text-green-600" />
<CircleCheck className="size-4 text-success" />
<span className="text-sm font-medium text-foreground">Upstream card loaded</span>
{card.version && <Badge variant="secondary">v{card.version}</Badge>}
{card.provider?.organization && <Badge variant="secondary">{card.provider.organization}</Badge>}

View file

@ -278,7 +278,7 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
Back to Agents
</Button>
<h1 className="text-2xl font-semibold">{agent.agent_name || "Unnamed Agent"}</h1>
<p className="text-sm text-gray-500 font-mono">{agent.agent_id}</p>
<p className="text-sm text-muted-foreground font-mono">{agent.agent_id}</p>
</div>
<Tabs value={activeTab} onValueChange={setActiveTab}>

View file

@ -16,8 +16,8 @@ const DocLink = ({ href, className }: DocLinkProps) => {
rel="noopener noreferrer"
title="Open documentation in a new tab"
className={cn(
"inline-flex items-center gap-2 rounded-xl border border-zinc-200 bg-white/80 px-3.5 py-2 text-sm font-medium text-zinc-700 shadow-xs",
"hover:bg-white focus-visible:outline-hidden focus-visible:ring-2 focus-visible:ring-blue-500 active:translate-y-[0.5px]",
"inline-flex items-center gap-2 rounded-xl border border-border bg-white/80 px-3.5 py-2 text-sm font-medium text-foreground shadow-xs",
"hover:bg-card focus-visible:outline-hidden focus-visible:ring-2 focus-visible:ring-ring active:translate-y-[0.5px]",
className,
)}
>

View file

@ -6,7 +6,7 @@
import { Plus, Wallet } from "lucide-react";
import React, { useCallback, useState } from "react";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { PageHeader } from "@/components/shared/PageHeader";
import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader";
import { ToolbarSeparator } from "@/components/shared/ToolbarSeparator";
import { Button } from "@/components/ui/button";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
@ -76,7 +76,7 @@ const BudgetPanel: React.FC<BudgetSettingsPageProps> = ({ accessToken }) => {
return (
<div className="flex h-full flex-col gap-4 p-6 px-12">
<PageHeader
<LegacyPageHeader
icon={<Wallet className="size-5" />}
title="Budgets"
subtitle="Spend, TPM and RPM limits you can assign to customers."

View file

@ -168,12 +168,12 @@ const HealthCheckDetails: React.FC<{ response: any }> = ({ response }) => {
<div>
<div className="mb-6 flex items-center">
{response?.status === "healthy" ? (
<CheckCircle2 className="mr-2 size-5 text-green-600" />
<CheckCircle2 className="mr-2 size-5 text-success" />
) : (
<XCircle className="mr-2 size-5 text-destructive" />
)}
<p
className={`text-sm font-medium ${response?.status === "healthy" ? "text-green-600" : "text-destructive"}`}
className={`text-sm font-medium ${response?.status === "healthy" ? "text-success" : "text-destructive"}`}
>
Cache Status: {response?.status || "unhealthy"}
</p>

View file

@ -107,7 +107,7 @@ const CoordinationRedisSettings: React.FC = () => {
manager. It is configured independently of the response cache.
</p>
<p className="text-xs text-muted-foreground">{badge.tooltip}</p>
<p className="text-xs text-amber-600 dark:text-amber-400">Saved changes take effect on proxy restart.</p>
<p className="text-xs text-warning">Saved changes take effect on proxy restart.</p>
</div>
<CoordinationRedisTypeSelector redisType={redisType} onTypeChange={setSelectedRedisType} />

View file

@ -8,7 +8,7 @@ export const ResponseTimeIndicator: React.FC<ResponseTimeIndicatorProps> = ({ re
if (responseTimeMs === null || responseTimeMs === undefined) return null;
return (
<div className="flex items-center space-x-1 text-xs text-gray-500 font-mono">
<div className="flex items-center space-x-1 text-xs text-muted-foreground font-mono">
<svg className="w-4 h-4" viewBox="0 0 24 24" fill="none" xmlns="http://www.w3.org/2000/svg">
<path
d="M12 6V12L16 14M12 2C6.47715 2 2 6.47715 2 12C2 17.5228 6.47715 22 12 22C17.5228 22 22 17.5228 22 12C22 6.47715 17.5228 2 12 2Z"

View file

@ -62,7 +62,7 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
<p className="text-5xl font-semibold tracking-tight text-foreground">{usd(stats.saved_spend)}</p>
<Badge
variant="secondary"
className={cheaper ? "bg-emerald-50 text-emerald-700" : "bg-red-50 text-destructive"}
className={cheaper ? "bg-success/10 text-success" : "bg-destructive/10 text-destructive"}
>
{stats.saved_spend !== 0 && (cheaper ? "-" : "+")}
{Math.abs(stats.saved_pct).toFixed(0)}%

View file

@ -40,7 +40,7 @@ const compareRows = (a: CacheLeakageRow, b: CacheLeakageRow, sort: SortState): n
const InfoTooltip = ({ info }: { info: string }) => (
<Tooltip>
<TooltipTrigger render={<span className="inline-flex" aria-label={info} />}>
<Info className="h-3 w-3 text-gray-400" />
<Info className="h-3 w-3 text-muted-foreground/70" />
</TooltipTrigger>
<TooltipContent className="max-w-xs">{info}</TooltipContent>
</Tooltip>
@ -72,7 +72,7 @@ const SortableHead = ({
className="inline-flex items-center gap-1 font-medium hover:text-foreground"
>
{label}
<Arrow className={`h-3 w-3 ${active ? "text-foreground" : "text-gray-400"}`} />
<Arrow className={`h-3 w-3 ${active ? "text-foreground" : "text-muted-foreground/70"}`} />
</button>
<InfoTooltip info={info} />
</span>

View file

@ -112,7 +112,7 @@ const PromptCompressionTab: React.FC<PromptCompressionTabProps> = ({ accessToken
href="https://docs.litellm.ai/docs/proxy/headroom"
target="_blank"
rel="noopener noreferrer"
className="text-blue-600 underline dark:text-blue-400"
className="text-info underline"
>
Headroom setup docs
</a>
@ -134,7 +134,7 @@ const PromptCompressionTab: React.FC<PromptCompressionTabProps> = ({ accessToken
<span
className={`rounded-full px-2 py-0.5 text-xs font-medium ${
guardrail.litellm_params?.default_on
? "bg-emerald-100 text-emerald-800 dark:bg-emerald-950 dark:text-emerald-300"
? "bg-success/15 text-success"
: "bg-muted text-muted-foreground"
}`}
>
@ -181,8 +181,8 @@ const PromptCompressionTab: React.FC<PromptCompressionTabProps> = ({ accessToken
)}
</FormField>
</FieldGroup>
<div className="mt-6 mb-4 rounded-lg border border-yellow-200 bg-yellow-50 p-3 dark:border-yellow-800 dark:bg-yellow-950">
<p className="text-sm text-yellow-800 dark:text-yellow-300">
<div className="mt-6 mb-4 rounded-lg border border-warning/20 bg-warning/10 p-3">
<p className="text-sm text-warning">
Applying compression to all requests is available to all users. Enabling it selectively per key or
team is a LiteLLM Enterprise feature. Get a trial key{" "}
<a

View file

@ -85,8 +85,8 @@ const endsIn = (endsAt: string | null | undefined): string | null => {
};
const STATUS_STYLES: Record<string, string> = {
running: "bg-blue-50 text-blue-700",
completed: "bg-emerald-50 text-emerald-700",
running: "bg-info/10 text-info",
completed: "bg-success/10 text-success",
stopped: "bg-secondary text-muted-foreground",
};
@ -146,8 +146,8 @@ const VerdictBar: React.FC<{ direction: ShadowEvalDirection; results: NonNullabl
? Math.max(0, 100 - results.overall_shadow_win_rate_pct - ties)
: results.overall_shadow_win_rate_pct;
const segments = [
{ label: "Router won", value: routerWins, fill: "bg-emerald-500" },
{ label: "Tie", value: ties, fill: "bg-emerald-200" },
{ label: "Router won", value: routerWins, fill: "bg-success" },
{ label: "Tie", value: ties, fill: "bg-success/20" },
{
label: `${otherArmLabel(direction)} won`,
value: Math.max(0, 100 - routerWins - ties),
@ -244,7 +244,7 @@ const JobResults: React.FC<{
)}
</div>
{(job.error_count ?? 0) > 0 && job.last_error != null && (
<p className="border-b bg-red-50 px-6 py-2 text-xs text-destructive">
<p className="border-b bg-destructive/10 px-6 py-2 text-xs text-destructive">
Last failure: <span className="font-mono">{job.last_error}</span>
</p>
)}

View file

@ -46,10 +46,10 @@ const SECTION_HEADER_CLASS = "group/section flex w-full items-center justify-bet
const SectionHeader: React.FC<{ title: string; description: string }> = ({ title, description }) => (
<CollapsibleTrigger className={SECTION_HEADER_CLASS}>
<div className="flex flex-col items-start w-full">
<span className="block text-lg font-semibold text-gray-900">{title}</span>
<span className="block text-sm text-gray-500 mt-1">{description}</span>
<span className="block text-lg font-semibold text-foreground">{title}</span>
<span className="block text-sm text-muted-foreground mt-1">{description}</span>
</div>
<ChevronDown className="size-5 shrink-0 text-gray-500 transition-transform group-data-[panel-open]/section:rotate-180" />
<ChevronDown className="size-5 shrink-0 text-muted-foreground transition-transform group-data-[panel-open]/section:rotate-180" />
</CollapsibleTrigger>
);
@ -177,17 +177,17 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
<div className="flex flex-col md:flex-row items-start md:items-center justify-between mb-6">
<div>
<div className="flex items-center gap-2">
<p className="text-xl font-medium text-gray-900">Cost Tracking Settings</p>
<p className="text-xl font-medium text-foreground">Cost Tracking Settings</p>
<DocsMenu items={DOCS_LINKS} />
</div>
<p className="text-gray-500 mt-1">
<p className="text-muted-foreground mt-1">
Configure cost discounts and margins for different LLM providers. Changes are saved automatically.
</p>
</div>
</div>
{/* Main Content Card with Accordions */}
<div className="bg-white rounded-lg shadow-sm w-full max-w-full space-y-4">
<div className="bg-card rounded-lg shadow-sm w-full max-w-full space-y-4">
{/* Accordion 1: Provider Discounts - Only for proxy admins */}
{isProxyAdmin && (
<Collapsible className="rounded-lg border">
@ -212,7 +212,7 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
</div>
{isFetching ? (
<div className="py-12 text-center">
<p className="text-gray-500">Loading configuration...</p>
<p className="text-muted-foreground">Loading configuration...</p>
</div>
) : Object.keys(discountConfig).length > 0 ? (
<ProviderDiscountTable
@ -223,7 +223,7 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
) : (
<div className="py-16 px-6 text-center">
<svg
className="mx-auto h-12 w-12 text-gray-400 mb-4"
className="mx-auto h-12 w-12 text-muted-foreground/70 mb-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
@ -235,8 +235,10 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
d="M12 8c-1.657 0-3 .895-3 2s1.343 2 3 2 3 .895 3 2-1.343 2-3 2m0-8c1.11 0 2.08.402 2.599 1M12 8V7m0 1v8m0 0v1m0-1c-1.11 0-2.08-.402-2.599-1M21 12a9 9 0 11-18 0 9 9 0 0118 0z"
/>
</svg>
<p className="text-gray-700 font-medium mb-2">No provider discounts configured</p>
<p className="text-gray-500 text-sm">Click &quot;Add Provider Discount&quot; to get started</p>
<p className="text-foreground font-medium mb-2">No provider discounts configured</p>
<p className="text-muted-foreground text-sm">
Click &quot;Add Provider Discount&quot; to get started
</p>
</div>
)}
</div>
@ -265,7 +267,7 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
</div>
{isFetching ? (
<div className="py-12 text-center">
<p className="text-gray-500">Loading configuration...</p>
<p className="text-muted-foreground">Loading configuration...</p>
</div>
) : Object.keys(marginConfig).length > 0 ? (
<ProviderMarginTable
@ -276,7 +278,7 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
) : (
<div className="py-16 px-6 text-center">
<svg
className="mx-auto h-12 w-12 text-gray-400 mb-4"
className="mx-auto h-12 w-12 text-muted-foreground/70 mb-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
@ -288,8 +290,10 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
d="M12 8c-1.657 0-3 .895-3 2s1.343 2 3 2 3 .895 3 2-1.343 2-3 2m0-8c1.11 0 2.08.402 2.599 1M12 8V7m0 1v8m0 0v1m0-1c-1.11 0-2.08-.402-2.599-1M21 12a9 9 0 11-18 0 9 9 0 0118 0z"
/>
</svg>
<p className="text-gray-700 font-medium mb-2">No provider margins configured</p>
<p className="text-gray-500 text-sm">Click &quot;Add Provider Margin&quot; to get started</p>
<p className="text-foreground font-medium mb-2">No provider margins configured</p>
<p className="text-muted-foreground text-sm">
Click &quot;Add Provider Margin&quot; to get started
</p>
</div>
)}
</div>
@ -334,12 +338,12 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
<Dialog open={isModalVisible} onOpenChange={(open) => !open && handleModalCancel()}>
<DialogContent className="top-8 max-h-[calc(100dvh-4rem)] translate-y-0 overflow-y-auto sm:max-w-[1000px]">
<DialogHeader>
<div className="flex items-center space-x-3 pb-4 border-b border-gray-100">
<DialogTitle className="text-xl font-semibold text-gray-900">Add Provider Discount</DialogTitle>
<div className="flex items-center space-x-3 pb-4 border-b border-border">
<DialogTitle className="text-xl font-semibold text-foreground">Add Provider Discount</DialogTitle>
</div>
</DialogHeader>
<div className="mt-6">
<p className="text-sm text-gray-600 mb-6">
<p className="text-sm text-muted-foreground mb-6">
Select a provider and set its discount percentage. Enter a value between 0% and 100% (e.g., 5 for a 5%
discount).
</p>
@ -360,12 +364,12 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
<Dialog open={isMarginModalVisible} onOpenChange={(open) => !open && handleMarginModalCancel()}>
<DialogContent className="top-8 max-h-[calc(100dvh-4rem)] translate-y-0 overflow-y-auto sm:max-w-[1000px]">
<DialogHeader>
<div className="flex items-center space-x-3 pb-4 border-b border-gray-100">
<DialogTitle className="text-xl font-semibold text-gray-900">Add Provider Margin</DialogTitle>
<div className="flex items-center space-x-3 pb-4 border-b border-border">
<DialogTitle className="text-xl font-semibold text-foreground">Add Provider Margin</DialogTitle>
</div>
</DialogHeader>
<div className="mt-6">
<p className="text-sm text-gray-600 mb-6">
<p className="text-sm text-muted-foreground mb-6">
Select a provider (or &quot;Global&quot; for all providers) and configure the margin. You can use
percentage-based or fixed amount.
</p>

Some files were not shown because too many files have changed in this diff Show more