mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
a1b32e446d
425 changed files with 9919 additions and 6952 deletions
|
|
@ -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
|
||||
;;
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
41
.github/actions/detect-changes/action.yml
vendored
Normal file
41
.github/actions/detect-changes/action.yml
vendored
Normal 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"
|
||||
|
|
@ -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}" ;;
|
||||
7
.github/workflows/_test-unit-base.yml
vendored
7
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
23
.github/workflows/test-linting.yml
vendored
23
.github/workflows/test-linting.yml
vendored
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
10
.github/workflows/test-litellm-ui-build.yml
vendored
10
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
11
.github/workflows/test-litellm-ui-unit.yml
vendored
11
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
9
.github/workflows/test-mcp.yml
vendored
9
.github/workflows/test-mcp.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
12
.github/workflows/test-unit-documentation.yml
vendored
12
.github/workflows/test-unit-documentation.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
|
|
@ -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 #############
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
|
|
|||
149
litellm/litellm_core_utils/ptu_pricing.py
Normal file
149
litellm/litellm_core_utils/ptu_pricing.py
Normal 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,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
456
litellm/proxy/hooks/batch_enqueued_tokens.py
Normal file
456
litellm/proxy/hooks/batch_enqueued_tokens.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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__ = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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", {
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 }) => {
|
||||
|
|
|
|||
|
|
@ -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 }) => {
|
||||
|
|
|
|||
163
tests/test_litellm/litellm_core_utils/test_ptu_pricing.py
Normal file
163
tests/test_litellm/litellm_core_utils/test_ptu_pricing.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
441
tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py
Normal file
441
tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py
Normal 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)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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"] == ""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
992
ui/litellm-dashboard/package-lock.json
generated
992
ui/litellm-dashboard/package-lock.json
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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 'Internal Users' 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 'Internal Users' 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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 & 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 "Require x-litellm-trace-id on calls BY this agent" 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'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'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">
|
||||
|
|
|
|||
|
|
@ -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>}
|
||||
|
|
|
|||
|
|
@ -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}>
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)}
|
||||
>
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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} />
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)}%
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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 "Add Provider Discount" 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 "Add Provider Discount" 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 "Add Provider Margin" 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 "Add Provider Margin" 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 "Global" 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
Loading…
Add table
Reference in a new issue