Merge litellm_internal_staging into MCP OAuth draft server branch

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-07-17 17:52:47 +00:00
commit a939a9f138
1969 changed files with 106947 additions and 24475 deletions

3
.github/CODEOWNERS vendored Normal file
View file

@ -0,0 +1,3 @@
/ui/ @yuneng-jiang @ryan-crabbe-berri
/litellm/proxy/_experimental/out/ @yuneng-jiang @ryan-crabbe-berri
/ui/litellm-dashboard/src/lib/http/schema.d.ts

View file

@ -0,0 +1,47 @@
name: "Set up uv with retries"
description: >-
Install uv via astral-sh/setup-uv, retrying on transient failures. Even with
an exact pinned version, the action resolves the artifact URL by fetching
https://raw.githubusercontent.com/astral-sh/versions/main/v1/uv.ndjson in a
single request with no retry, timeout, or fallback, so one connection-level
network error ("fetch failed") fails the whole job before any test runs.
Retrying the full step covers the manifest fetch and the binary download.
inputs:
version:
description: "uv version to install"
required: true
runs:
using: composite
steps:
- name: Set up uv (attempt 1)
id: attempt-1
continue-on-error: true
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
with:
version: ${{ inputs.version }}
- name: Wait before attempt 2
if: steps.attempt-1.outcome == 'failure'
shell: bash
run: sleep 15
- name: Set up uv (attempt 2)
id: attempt-2
if: steps.attempt-1.outcome == 'failure'
continue-on-error: true
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
with:
version: ${{ inputs.version }}
- name: Wait before attempt 3
if: steps.attempt-2.outcome == 'failure'
shell: bash
run: sleep 30
- name: Set up uv (attempt 3)
if: steps.attempt-2.outcome == 'failure'
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
with:
version: ${{ inputs.version }}

View file

@ -41,3 +41,27 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
✅ Test
## Changes
## QA runbook
<!-- Only needed when your PR edits tests/e2e; delete this section otherwise
For each e2e test you added or changed, list the manual steps a reviewer can follow to reproduce it by hand against a live proxy, mapping 1:1 to what the test asserts: one top-level bullet per test giving its pytest node id followed by what it proves in plain words, then a nested "- [ ]" checklist where each item is a concrete action (route, request body, expected response) and the final item is the sanity-check step shown in the examples. Note environment prerequisites (provider credentials, config flags) and any nuances a manual run will hit. See PRs #32914 and #32963 for full examples
Example checklists:
- tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py::TestKeyRateLimits::test_rpm_limit_blocks_over_limit - a key allowed 2 requests a minute serves exactly 2 and refuses the 3rd
- [ ] Generate a limited key: curl -X POST http://localhost:4000/key/generate -H "Authorization: Bearer sk-1234" -d '{"rpm_limit": 2}'
- [ ] Send three /v1/chat/completions requests with that key inside one minute
- [ ] Expect the first two to return 200 and the third to return 429 naming the rpm limit
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
- tests/e2e/management/test_management_e2e.py::TestModelRoutes::test_model_create_appears_in_ui - a deployment created through the API shows up on the Admin UI models page
- [ ] POST /model/new with the master key, a bedrock model, and aws_region_name (needs STORE_MODEL_IN_DB=True and AWS credentials)
- [ ] Open http://localhost:4000/ui/?page=models and expect a deployment row showing the returned model id
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
-->
### Final Attestation
- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR

View file

@ -63,7 +63,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -18,7 +18,7 @@ jobs:
with:
persist-credentials: false
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Update JSON Data

View file

@ -31,7 +31,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -37,7 +37,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -31,12 +31,12 @@ jobs:
echo "PR head repo: $HEAD_REPO"
echo "PR head branch: $HEAD_REF"
if [ "$HEAD_REPO" != "$BASE_REPO" ]; then
echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against the 'litellm_oss_staging' branch instead."
echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead."
exit 1
fi
if [ "$HEAD_REF" = "litellm_internal_staging" ] || [[ "$HEAD_REF" == litellm_hotfix_?* ]]; then
echo "Allowed source branch."
exit 0
fi
echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against 'litellm_oss_staging' instead."
echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead."
exit 1

View file

@ -39,7 +39,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -35,7 +35,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -38,7 +38,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -33,7 +33,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
@ -48,7 +48,7 @@ jobs:
- name: Install dependencies
run: |
uv sync --frozen --group proxy-dev
uv sync --frozen --group proxy-dev --group e2e-dev
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
@ -107,6 +107,16 @@ jobs:
run: |
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
- name: Check tests/e2e basedpyright (zero errors)
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
if git diff --name-only --diff-filter=ACMRD "$BASE_SHA"...HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
uv run --no-sync basedpyright tests/e2e
else
echo "No changed tests/e2e Python files; skipping."
fi
- name: Check for circular imports
run: |
cd litellm
@ -162,7 +172,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -36,79 +36,3 @@ jobs:
- name: Build
run: npm run build
frontend-lint:
runs-on: ubuntu-latest
timeout-minutes: 8
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
persist-credentials: false
- name: Collect changed files
id: changed
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
[ -f "$f" ] || continue
case "$f" in
*.js | *.jsx | *.ts | *.tsx | *.mjs | *.cjs)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt"
printf '%s\n' "$f" >> "$RUNNER_TEMP/eslint_files.txt" ;;
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$BASE_SHA"...HEAD -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else
echo "has_files=false" >> "$GITHUB_OUTPUT"
echo "No lintable UI files changed in this PR; nothing to check."
fi
- name: Setup Node.js
if: steps.changed.outputs.has_files == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version: "20"
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changed.outputs.has_files == 'true'
run: npm ci
- name: Lint changed files (prettier + eslint)
if: steps.changed.outputs.has_files == 'true'
run: |
prettier_files=()
eslint_files=()
while IFS= read -r f; do prettier_files+=("$f"); done < "$RUNNER_TEMP/prettier_files.txt"
while IFS= read -r f; do eslint_files+=("$f"); done < "$RUNNER_TEMP/eslint_files.txt"
status=0
if [ ${#prettier_files[@]} -gt 0 ]; then
echo "::group::Prettier (${#prettier_files[@]} files)"
npx prettier --check "${prettier_files[@]}" || { status=1; echo "::error::Unformatted files. Fix with: npm run format"; }
echo "::endgroup::"
fi
if [ ${#eslint_files[@]} -gt 0 ]; then
echo "::group::ESLint (${#eslint_files[@]} files)"
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_files[@]}" || status=1
echo "::endgroup::"
fi
exit $status
- name: Check lint budgets
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: |
npx eslint . -f json -o "$RUNNER_TEMP/lint-report.json" || true
node scripts/check-lint-budgets.mjs "$RUNNER_TEMP/lint-report.json" eslint-budgets.json --check eslint-metrics.json

View file

@ -0,0 +1,92 @@
name: UI Lint
permissions:
contents: read
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
jobs:
frontend-lint:
runs-on: ubuntu-latest
timeout-minutes: 8
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
persist-credentials: false
- name: Collect changed files
id: changed
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
[ -f "$f" ] || continue
case "$f" in
*.js | *.jsx | *.ts | *.tsx | *.mjs | *.cjs)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt"
printf '%s\n' "$f" >> "$RUNNER_TEMP/eslint_files.txt" ;;
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$BASE_SHA"...HEAD -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else
echo "has_files=false" >> "$GITHUB_OUTPUT"
echo "No lintable UI files changed in this PR; nothing to check."
fi
- name: Setup Node.js
if: steps.changed.outputs.has_files == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version: "20"
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changed.outputs.has_files == 'true'
run: npm ci
- name: Lint changed files (prettier + eslint)
if: steps.changed.outputs.has_files == 'true'
run: |
prettier_files=()
eslint_files=()
while IFS= read -r f; do prettier_files+=("$f"); done < "$RUNNER_TEMP/prettier_files.txt"
while IFS= read -r f; do eslint_files+=("$f"); done < "$RUNNER_TEMP/eslint_files.txt"
status=0
if [ ${#prettier_files[@]} -gt 0 ]; then
echo "::group::Prettier (${#prettier_files[@]} files)"
npx prettier --check "${prettier_files[@]}" || { status=1; echo "::error::Unformatted files. Fix with: npm run format"; }
echo "::endgroup::"
fi
if [ ${#eslint_files[@]} -gt 0 ]; then
echo "::group::ESLint (${#eslint_files[@]} files)"
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_files[@]}" || status=1
echo "::endgroup::"
fi
exit $status
- name: Check lint budgets
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: |
npx eslint . -f json -o "$RUNNER_TEMP/lint-report.json" || true
node scripts/check-lint-budgets.mjs "$RUNNER_TEMP/lint-report.json" eslint-budgets.json
- name: Check for dead code (knip)
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: npm run knip:ci

View file

@ -32,7 +32,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -31,7 +31,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -74,7 +74,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -42,7 +42,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -5,6 +5,8 @@ on:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:
contents: read

View file

@ -59,7 +59,7 @@ jobs:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"

View file

@ -4,7 +4,11 @@ on:
push:
branches: [main, litellm_internal_staging]
pull_request:
branches: [main, litellm_internal_staging]
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}

7
.gitignore vendored
View file

@ -106,6 +106,13 @@ STABILIZATION_TODO.md
**/coverage
test-config
# Claude Code compatibility-matrix pytest artifact (CI-only output).
compat-results.json
compat-results.json.shards/
compat-rate-limit-summary.json
# Matrix JSON produced by the daily-cron publisher (pushed to litellm-docs).
compatibility-matrix.json
# ---------- Terraform ----------
# Provider binaries + module cache — regenerated by `terraform init`.
**/.terraform/

View file

@ -19,7 +19,7 @@ Same thing for bug fixes. The tests should make it so that this specific bug can
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose for internal contributors; external / OSS contributions target the current daily OSS branch instead, named `litellm_oss_daily_YYYY_MM_DD` (a fresh one is cut each weekday, so use the most recent)
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
@ -39,6 +39,8 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a
Python max line length is 120, not 88
On a fresh worktree or clone, run `make bootstrap` before anything else. It provisions everything tests, `make pre-commit`, and a local proxy need
Run tests before you commit. Also, run `make pre-commit` right before each commit, which generates types (as needed) and formats/lints your code. Any errors found must be fixed. It only runs when there are staged frontend and/or backend changes and calculates violations, generates types, etc. based on the worktree, so stage what you need or stash/delete unwanted files in litellm/ or ui/ (where backend and frontend lint run, respectively) before running it. If it fails because dashboard api types are stale, it already regenerated them for you. You just need to stage the schema.d.ts, re-run `make pre-commit` to confirm it passes, and commit
When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered limits so the ceilings ratchet down instead of leaving stale headroom. It measures the working tree, so it must contain exactly the fixes you're committing

View file

@ -322,7 +322,7 @@ npm run build
## Submitting Your PR
1. **Push your branch**: `git push origin your-feature-branch`
2. **Create a PR**: Go to GitHub and create a pull request
2. **Create a PR**: Go to GitHub and open a pull request against the current daily OSS branch, named `litellm_oss_daily_YYYY_MM_DD`. A fresh one is cut each weekday, so pick the most recent from the [branch list](https://github.com/BerriAI/litellm/branches/all?query=litellm_oss_daily). Do not target `main`.
3. **Fill out the PR template**: Provide clear description of changes
4. **Wait for review**: Maintainers will review and provide feedback
5. **Address feedback**: Make requested changes and push updates

View file

@ -114,6 +114,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy only the Prisma subdirs — copying the
# whole /root/.cache drags in the uv build cache (~660 MB, includes a

View file

@ -5,15 +5,16 @@
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev lint-checks format \
lint-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety pre-commit \
lint-install lint-fetch-base
lint-install lint-fetch-base bootstrap
# Default target
help:
@echo "Available commands:"
@echo " make bootstrap - Provision a fresh clone/worktree"
@echo " make install-dev - Install development dependencies"
@echo " make install-proxy-dev - Install proxy development dependencies"
@echo " make install-dev-ci - Install dev dependencies (CI-compatible, pins OpenAI)"
@ -27,6 +28,7 @@ help:
@echo " make lint - Run all linting (Ruff, basedpyright, format check, circular imports, import safety)"
@echo " make lint-ruff - Run Ruff linting only"
@echo " make lint-basedpyright - Run basedpyright strict, gated by per-rule error counts"
@echo " make lint-e2e-basedpyright - Run basedpyright over tests/e2e (zero errors allowed)"
@echo " make lint-basedpyright-budget-update - Ratchet basedpyright limits down by what this branch fixed"
@echo " make lint-format - Check ruff format formatting (matches CI)"
@echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its limit"
@ -54,6 +56,7 @@ UV := uv
UV_RUN := $(UV) run --no-sync
LINT_DEP_INSTALL ?= install-dev
LINT_E2E_DEP_INSTALL ?= lint-install
LINT_DEP_BASE ?= lint-fetch-base
LINT_JOBS := $(shell sysctl -n hw.ncpu 2>/dev/null || nproc 2>/dev/null || echo 4)
LINT_OUTPUT_SYNC := $(if $(filter output-sync,$(.FEATURES)),--output-sync=target,)
@ -69,6 +72,18 @@ info:
install-dev:
$(UV) sync --inexact --frozen
bootstrap:
$(UV) sync --inexact --frozen --extra proxy --group proxy-dev --group e2e-dev
$(UV_RUN) python scripts/prisma_generate_if_needed.py
cd ui/litellm-dashboard && npm ci --no-audit --no-fund
@main_root=$$(git worktree list --porcelain | head -1 | sed 's/^worktree //'); \
if [ "$$main_root" != "$$(git rev-parse --show-toplevel)" ] && [ -f "$$main_root/.env" ] && [ ! -f .env ]; then \
cp "$$main_root/.env" .env && echo "bootstrap: copied .env from $$main_root"; \
else \
echo "bootstrap: .env left untouched"; \
fi
@echo "bootstrap: done"
install-proxy-dev:
$(UV) sync --frozen --group proxy-dev --extra proxy
@ -111,7 +126,7 @@ lint-fetch-base:
# CI's). --inexact tops up the venv instead of pruning the proxy extras gen:api and the
# running proxy need.
lint-install:
$(UV) sync --inexact --frozen --group proxy-dev
$(UV) sync --inexact --frozen --group proxy-dev --group e2e-dev
$(UV_RUN) python scripts/prisma_generate_if_needed.py
# Diff-scoped format check, identical to test-linting.yml's "Check ruff format" step:
@ -164,6 +179,9 @@ lint-ruff-FULL-dev: install-dev
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
lint-e2e-basedpyright: $(LINT_E2E_DEP_INSTALL)
$(UV_RUN) basedpyright tests/e2e
# Type-discipline budget (mutable collections / casts / type guards / kwargs /
# unexplained suppressions), the test-linting.yml step `make lint` used to omit.
lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
@ -208,9 +226,9 @@ check-import-safety: $(LINT_DEP_INSTALL)
# base fetch) runs once up front; the checks themselves are independent, so a sub-make
# fans them out with -j and the fast ones finish under basedpyright's shadow.
lint: lint-install lint-fetch-base
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_DEP_BASE= lint-checks
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright check-circular-imports check-import-safety
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety
# Faster linting for local development (only checks changed code)
lint-dev: lint-format-changed check-circular-imports check-import-safety

View file

@ -552,17 +552,12 @@ The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws
2. Run dependent services `docker-compose up db prometheus`
#### Backend
1. (In root) create virtual environment `python -m venv .venv`
2. Activate virtual environment `source .venv/bin/activate`
3. Install dependencies `uv sync --all-extras --group proxy-dev`
4. `uv run prisma generate`
5. `prisma generate`
6. Start proxy backend `python litellm/proxy/proxy_cli.py`
1. Run `make bootstrap`
2. Start proxy backend: `uv run python litellm/proxy/proxy_cli.py`
#### Frontend
1. Navigate to `ui/litellm-dashboard`
2. Install dependencies `npm install`
3. Run `npm run dev` to start the dashboard
1. Navigate to `ui/litellm-dashboard` (dependencies were already installed w/ `make bootstrap`)
2. Start dashboard: `npm run dev`
### Verify Docker Image Signatures

View file

@ -46,6 +46,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/fallback",
"/fallbacks",
"/cache_settings",
"/coordination_redis/",
"/cost_tracking",
"/cost/",
"/credentials",

View file

@ -57,7 +57,7 @@
"limit": 5900
},
"reportMissingTypeArgument": {
"limit": 15918
"limit": 15903
},
"reportMissingTypeStubs": {
"limit": 41
@ -105,13 +105,13 @@
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40541
"limit": 40539
},
"reportUnknownParameterType": {
"limit": 20418
"limit": 20403
},
"reportUnknownVariableType": {
"limit": 32150
"limit": 32141
},
"reportUnnecessaryCast": {
"limit": 177
@ -123,7 +123,7 @@
"limit": 7
},
"reportUnnecessaryIsInstance": {
"limit": 1211
"limit": 1209
},
"reportUntypedBaseClass": {
"limit": 165

View file

@ -111,6 +111,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem

View file

@ -137,6 +137,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
COPY --from=builder /app/.cache /app/.cache
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets

View file

@ -36,7 +36,7 @@ RUN uv venv --python python && \
"opentelemetry-api==1.28.0" \
"opentelemetry-sdk==1.28.0" \
"opentelemetry-exporter-otlp==1.28.0" \
"ddtrace==2.19.0" \
"ddtrace==4.11.0" \
"sentry-sdk==2.21.0" \
"mangum==0.17.0" \
"azure-ai-contentsafety==1.0.0" \

View file

@ -7,7 +7,8 @@
# Thank you users! We ❤️ you! - Krrish & Ishaan
## This provides an LLM Guard Integration for content moderation on the proxy
from typing import Literal, Optional
import asyncio
from typing import Optional
import aiohttp
from fastapi import HTTPException
@ -18,7 +19,6 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import CallTypesLiteral
from litellm.utils import get_formatted_prompt
class _ENTERPRISE_LLMGuard(CustomLogger):
@ -46,45 +46,44 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
except Exception:
pass
async def moderation_check(self, text: str):
async def moderation_check(self, text: str) -> str:
"""
Runs the LLM Guard moderation check on ``text``.
Raises an HTTPException when the content violates the safety policy;
otherwise returns the sanitized prompt from LLM Guard, falling back to
the original text when the API does not provide one.
[TODO] make this more performant for high-throughput scenario
"""
try:
async with aiohttp.ClientSession() as session:
if self.mock_redacted_text is not None:
redacted_text = self.mock_redacted_text
else:
# Make the first request to /analyze
analyze_url = f"{self.llm_guard_api_base}analyze/prompt"
verbose_proxy_logger.debug("Making request to: %s", analyze_url)
analyze_payload = {"prompt": text}
redacted_text = None
if self.mock_redacted_text is not None:
redacted_text = self.mock_redacted_text
else:
analyze_url = f"{self.llm_guard_api_base}analyze/prompt"
verbose_proxy_logger.debug("Making request to: %s", analyze_url)
async with aiohttp.ClientSession() as session:
async with session.post(
analyze_url, json=analyze_payload
analyze_url, json={"prompt": text}
) as response:
redacted_text = await response.json()
verbose_proxy_logger.debug(
f"LLM Guard: Received response - {redacted_text}"
verbose_proxy_logger.debug(
f"LLM Guard: Received response - {redacted_text}"
)
if redacted_text is None:
raise HTTPException(
status_code=500,
detail={
"error": f"Invalid content moderation response: {redacted_text}"
},
)
if redacted_text is not None:
if (
redacted_text.get("is_valid", None) is not None
and redacted_text["is_valid"] is False
):
raise HTTPException(
status_code=400,
detail={"error": "Violated content safety policy"},
)
else:
pass
else:
raise HTTPException(
status_code=500,
detail={
"error": f"Invalid content moderation response: {redacted_text}"
},
)
if redacted_text.get("is_valid", None) is False:
raise HTTPException(
status_code=400,
detail={"error": "Violated content safety policy"},
)
sanitized_prompt = redacted_text.get("sanitized_prompt")
return sanitized_prompt if isinstance(sanitized_prompt, str) else text
except Exception as e:
verbose_proxy_logger.exception(
"litellm.enterprise.enterprise_hooks.llm_guard::moderation_check - Exception occurred - {}".format(
@ -138,23 +137,75 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
return
self.print_verbose("Makes LLM Guard Check")
try:
assert call_type in [
"completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
]
except Exception:
if call_type not in [
"completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
]:
self.print_verbose(
f"Call Type - {call_type}, not in accepted list - ['completion','embeddings','image_generation','moderation','audio_transcription']"
)
return data
formatted_prompt = get_formatted_prompt(data=data, call_type=call_type) # type: ignore
self.print_verbose(f"LLM Guard, formatted_prompt: {formatted_prompt}")
return await self.moderation_check(text=formatted_prompt)
return await self._moderate_request(data=data)
async def _moderate_request(self, data: dict) -> dict:
"""
Sanitizes the request in place using the prompt returned by LLM Guard so
the provider-bound request carries the redacted content, then returns it.
"""
messages = data.get("messages")
if messages is not None:
data["messages"] = list(
await asyncio.gather(
*(self._moderate_message(message) for message in messages)
)
)
return data
input_ = data.get("input")
if input_ is not None:
data["input"] = await self._moderate_input(input_)
return data
prompt = data.get("prompt")
if isinstance(prompt, str):
data["prompt"] = await self.moderation_check(text=prompt)
return data
async def _moderate_message(self, message: dict) -> dict:
content = message.get("content")
if isinstance(content, str):
return {**message, "content": await self.moderation_check(text=content)}
if isinstance(content, list):
return {
**message,
"content": list(
await asyncio.gather(
*(self._moderate_content_part(part) for part in content)
)
),
}
return message
async def _moderate_content_part(self, part: dict) -> dict:
if part.get("type") == "text" and isinstance(part.get("text"), str):
return {**part, "text": await self.moderation_check(text=part["text"])}
return part
async def _moderate_input(self, input_: object) -> object:
if isinstance(input_, str):
return await self.moderation_check(text=input_)
if isinstance(input_, list):
return [
await self.moderation_check(text=item)
if isinstance(item, str)
else item
for item in input_
]
return input_
async def async_post_call_streaming_hook(
self, user_api_key_dict: UserAPIKeyAuth, response: str

View file

@ -113,6 +113,10 @@ class PagerDutyAlerting(SlackAlerting):
user_api_key_spend=_meta.get("user_api_key_spend"),
user_api_key_max_budget=_meta.get("user_api_key_max_budget"),
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
user_api_key_user_spend=_meta.get("user_api_key_user_spend"),
user_api_key_user_max_budget=_meta.get("user_api_key_user_max_budget"),
user_api_key_team_spend=_meta.get("user_api_key_team_spend"),
user_api_key_team_max_budget=_meta.get("user_api_key_team_max_budget"),
user_api_key_org_id=_meta.get("user_api_key_org_id"),
user_api_key_org_alias=_meta.get("user_api_key_org_alias"),
user_api_key_team_id=_meta.get("user_api_key_team_id"),
@ -196,6 +200,10 @@ class PagerDutyAlerting(SlackAlerting):
if user_api_key_dict.budget_reset_at
else None
),
user_api_key_user_spend=user_api_key_dict.user_spend,
user_api_key_user_max_budget=user_api_key_dict.user_max_budget,
user_api_key_team_spend=user_api_key_dict.team_spend,
user_api_key_team_max_budget=user_api_key_dict.team_max_budget,
user_api_key_org_id=user_api_key_dict.org_id,
user_api_key_org_alias=user_api_key_dict.organization_alias,
user_api_key_team_id=user_api_key_dict.team_id,

View file

@ -29,7 +29,7 @@ class CheckBatchCost:
proxy_logging_obj: "ProxyLogging",
prisma_client: "PrismaClient",
llm_router: "Router",
track_unmanaged_vertex_batch_cost: bool = False,
track_unmanaged_batch_cost: bool = False,
):
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router import Router
@ -37,7 +37,7 @@ class CheckBatchCost:
self.proxy_logging_obj: ProxyLogging = proxy_logging_obj
self.prisma_client: PrismaClient = prisma_client
self.llm_router: Router = llm_router
self._track_unmanaged_vertex_batch_cost = track_unmanaged_vertex_batch_cost
self._track_unmanaged_batch_cost = track_unmanaged_batch_cost
# Cached after the first poll cycle. Once we know the column is absent we skip
# the guaranteed-failing primary query on every subsequent cycle.
self._has_batch_processed_column: bool = True
@ -118,11 +118,11 @@ class CheckBatchCost:
Resolve (model_id, batch_id) for a managed-object row, where model_id is a router
deployment id and batch_id is the raw provider batch id.
Managed batches encode both in a base64 unified id. Unmanaged Vertex batches, created with
a raw gs:// input_file_id, store the raw provider job id as unified_object_id; when
track_unmanaged_vertex_batch_cost is enabled the model is derived from the gs:// path and
mapped to a configured vertex_ai deployment. Returns None (recording a metric) when the row
can't be routed.
Managed batches encode both in a base64 unified id. Unmanaged batches (created outside
LiteLLM's own /v1/batches with a raw input_file_id) store the raw provider job id as
unified_object_id instead; when track_unmanaged_batch_cost is enabled the model is derived
from the provider-specific input_file_id layout (Vertex gs:// or Bedrock s3://) and mapped
to a matching deployment. Returns None (recording a metric) when the row can't be routed.
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
@ -142,8 +142,43 @@ class CheckBatchCost:
return None
return model_id, get_batch_id_from_unified_batch_id(decoded)
if self._track_unmanaged_vertex_batch_cost:
return self._resolve_unmanaged_vertex_routing(job, prom_logger)
if self._track_unmanaged_batch_cost:
from litellm.llms.bedrock.batches.transformation import (
BedrockBatchesConfig,
)
from litellm.llms.vertex_ai.batches.transformation import (
VertexAIBatchTransformation,
)
input_file_id = self._get_input_file_id(job)
if VertexAIBatchTransformation.is_unmanaged_gcs_batch_input_file_id(
input_file_id
):
assert input_file_id is not None # narrowed by is_unmanaged_gcs_batch_input_file_id
return self._resolve_unmanaged_provider_routing(
job=job,
prom_logger=prom_logger,
llm_provider="vertex_ai",
bare_model_name=VertexAIBatchTransformation.get_bare_model_name_from_gcs_file(
input_file_id
),
)
if BedrockBatchesConfig.is_unmanaged_s3_batch_input_file_id(input_file_id):
assert input_file_id is not None # narrowed by is_unmanaged_s3_batch_input_file_id
return self._resolve_unmanaged_provider_routing(
job=job,
prom_logger=prom_logger,
llm_provider="bedrock",
bare_model_name=BedrockBatchesConfig.get_bare_model_name_from_s3_file(
input_file_id
),
)
verbose_proxy_logger.info(
f"Skipping job {unified_object_id}: not a recognized unmanaged batch "
"(no gs:// or s3:// input_file_id with an embedded model)"
)
self._record_error(prom_logger, "invalid_unified_id")
return None
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} because it is not a valid unified object id"
@ -151,36 +186,17 @@ class CheckBatchCost:
self._record_error(prom_logger, "invalid_unified_id")
return None
def _resolve_unmanaged_vertex_routing(
def _resolve_unmanaged_provider_routing(
self,
job: "LiteLLM_ManagedObjectTable",
prom_logger: Optional["PrometheusLogger"],
llm_provider: str,
bare_model_name: str,
) -> Optional[Tuple[str, str]]:
from litellm.llms.vertex_ai.batches.transformation import (
VertexAIBatchTransformation,
)
input_file_id = self._get_input_file_id(job)
if not VertexAIBatchTransformation.is_unmanaged_gcs_batch_input_file_id(
input_file_id
):
verbose_proxy_logger.info(
f"Skipping job {job.unified_object_id}: not an unmanaged vertex batch "
"(no gs:// input_file_id with a publishers/ model path)"
)
self._record_error(prom_logger, "invalid_unified_id")
return None
assert input_file_id is not None # narrowed by is_unmanaged_gcs_batch_input_file_id
bare_model_name = VertexAIBatchTransformation.get_bare_model_name_from_gcs_file(
input_file_id
)
deployment_id = self._get_vertex_ai_deployment_id_for_bare_model(
bare_model_name
)
deployment_id = self._get_deployment_id_for_bare_model(bare_model_name, llm_provider)
if deployment_id is None:
verbose_proxy_logger.info(
f"Skipping unmanaged vertex batch {job.unified_object_id}: no vertex_ai "
f"Skipping unmanaged {llm_provider} batch {job.unified_object_id}: no {llm_provider} "
f"deployment configured for model {bare_model_name}"
)
self._record_error(prom_logger, "unmanaged_no_matching_deployment")
@ -188,22 +204,22 @@ class CheckBatchCost:
return deployment_id, job.unified_object_id
def _get_vertex_ai_deployment_id_for_bare_model(
self, bare_model_name: str
def _get_deployment_id_for_bare_model(
self, bare_model_name: str, llm_provider: str
) -> Optional[str]:
model_group = self.llm_router.resolve_model_name_from_model_id(bare_model_name)
deployment_id = (
self._get_vertex_ai_deployment_id(model_group) if model_group else None
self._get_deployment_id_for_provider(model_group, llm_provider) if model_group else None
)
if deployment_id is not None:
return deployment_id
return self._get_vertex_ai_deployment_id_from_matching_deployments(
bare_model_name
return self._get_deployment_id_from_matching_deployments(
bare_model_name, llm_provider
)
def _get_vertex_ai_deployment_id_from_matching_deployments(
self, bare_model_name: str
def _get_deployment_id_from_matching_deployments(
self, bare_model_name: str, llm_provider: str
) -> Optional[str]:
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@ -215,13 +231,13 @@ class CheckBatchCost:
if not self._is_bare_model_match(actual_model, bare_model_name):
continue
try:
_, llm_provider, _, _ = get_llm_provider(
_, deployment_llm_provider, _, _ = get_llm_provider(
model=actual_model,
custom_llm_provider=litellm_params.get("custom_llm_provider"),
)
except Exception:
continue
if llm_provider != "vertex_ai":
if deployment_llm_provider != llm_provider:
continue
model_info = deployment.get("model_info") or {}
deployment_id = model_info.get("id")
@ -231,15 +247,21 @@ class CheckBatchCost:
@staticmethod
def _is_bare_model_match(actual_model: str, bare_model_name: str) -> bool:
# Bedrock model ids may have ":" replaced with "-" in the S3 object key (see
# BedrockBatchesConfig.get_bare_model_name_from_s3_file), so normalize both sides;
# a no-op for providers like vertex_ai whose model ids never contain a colon.
normalized_actual = actual_model.replace(":", "-")
normalized_bare = bare_model_name.replace(":", "-")
return (
actual_model == bare_model_name
or actual_model.endswith(f"/{bare_model_name}")
or actual_model.endswith(f":{bare_model_name}")
normalized_actual == normalized_bare
or normalized_actual.endswith(f"/{normalized_bare}")
)
def _get_vertex_ai_deployment_id(self, model_group: str) -> Optional[str]:
def _get_deployment_id_for_provider(
self, model_group: str, llm_provider: str
) -> Optional[str]:
"""
Returns the first deployment id for `model_group` whose provider is vertex_ai,
Returns the first deployment id for `model_group` whose provider is `llm_provider`,
skipping deployments from other providers that happen to share the model group name.
"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@ -249,13 +271,13 @@ class CheckBatchCost:
if deployment_info is None:
continue
try:
_, llm_provider, _, _ = get_llm_provider(
_, deployment_llm_provider, _, _ = get_llm_provider(
model=deployment_info.litellm_params.model,
custom_llm_provider=deployment_info.litellm_params.custom_llm_provider,
)
except Exception:
continue
if llm_provider == "vertex_ai":
if deployment_llm_provider == llm_provider:
return deployment_id
return None

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.49"
version = "0.1.51"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.49"
version = "0.1.51"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -54,6 +54,12 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `pdb.annotations` | Extra metadata annotations to add to the PDB | `{}` |
| `pdb.labels` | Extra metadata labels to add to the PDB | `{}` |
| `billingMetrics.enabled` | Enable enterprise billable-request metering. Requires an enterprise license. | `false` |
| `billingMetrics.endpoint` | Collector that the billable-request counter is pushed to. | `https://telemetry.litellm.ai` |
| `billingMetrics.secretName` | Name of an existing Secret holding the mTLS client certificate, under the keys `tls.crt` and `tls.key`. | `litellm-billing-metrics-mtls` |
| `billingMetrics.caSecretName` | Name of an existing Secret holding a CA bundle under the key `ca.crt`. Only needed for a private or test collector whose server certificate is not on the public web PKI. | `""` |
| `billingMetrics.exportIntervalMs` | How often the counter is pushed, in milliseconds. The proxy defaults to `60000` when unset. | `""` |
#### Example `proxy_config` ConfigMap from values (default):
```
@ -94,6 +100,21 @@ data:
type: Opaque
```
#### Enterprise billable-request metering
Enterprise licenses meter billable requests by pushing a counter to LiteLLM's collector over mutual TLS. The chart does not create the client certificate; it mounts one you already hold, read-only, so the private key is never exposed through the environment. Create the Secret under the name the chart expects, then turn the block on:
```
kubectl create secret tls litellm-billing-metrics-mtls --cert=client.crt --key=client.key
```
```
billingMetrics:
enabled: true
```
Set `billingMetrics.caSecretName` only when the collector is a private or test one whose server certificate is not on the public web PKI; the production collector needs no CA override. The chart fails the render rather than deploying a proxy that silently never exports, so a missing `secretName` or an emptied `endpoint` surfaces at `helm install` time.
### Database Settings
| Name | Description | Value |

View file

@ -50,6 +50,53 @@ app.kubernetes.io/name: {{ include "litellm.name" . }}
app.kubernetes.io/instance: {{ .Release.Name }}
{{- end }}
{{/*
Enterprise billable-request metering. The client certificate identifies the
deployment to LiteLLM's collector, so it is mounted read-only from an existing
Secret rather than passed through the environment.
*/}}
{{- define "litellm.billingMetrics.certDir" -}}/etc/litellm/billing-mtls{{- end -}}
{{- define "litellm.billingMetrics.caDir" -}}/etc/litellm/billing-mtls-ca{{- end -}}
{{- define "litellm.billingMetricsEnv" -}}
- name: LITELLM_BILLING_METRICS_ENDPOINT
value: {{ required "billingMetrics.endpoint is required when billingMetrics.enabled is true" .Values.billingMetrics.endpoint | quote }}
- name: LITELLM_BILLING_METRICS_CLIENT_CERT
value: {{ printf "%s/tls.crt" (include "litellm.billingMetrics.certDir" .) | quote }}
- name: LITELLM_BILLING_METRICS_CLIENT_KEY
value: {{ printf "%s/tls.key" (include "litellm.billingMetrics.certDir" .) | quote }}
{{- if .Values.billingMetrics.caSecretName }}
- name: LITELLM_BILLING_METRICS_CA_CERT
value: {{ printf "%s/ca.crt" (include "litellm.billingMetrics.caDir" .) | quote }}
{{- end }}
{{- with .Values.billingMetrics.exportIntervalMs }}
- name: LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS
value: {{ . | quote }}
{{- end }}
{{- end -}}
{{- define "litellm.billingMetricsVolumes" -}}
- name: billing-metrics-mtls
secret:
secretName: {{ required "billingMetrics.secretName is required when billingMetrics.enabled is true (an existing Secret with tls.crt and tls.key)" .Values.billingMetrics.secretName }}
{{- if .Values.billingMetrics.caSecretName }}
- name: billing-metrics-mtls-ca
secret:
secretName: {{ .Values.billingMetrics.caSecretName }}
{{- end }}
{{- end -}}
{{- define "litellm.billingMetricsVolumeMounts" -}}
- name: billing-metrics-mtls
mountPath: {{ include "litellm.billingMetrics.certDir" . }}
readOnly: true
{{- if .Values.billingMetrics.caSecretName }}
- name: billing-metrics-mtls-ca
mountPath: {{ include "litellm.billingMetrics.caDir" . }}
readOnly: true
{{- end }}
{{- end -}}
{{/*
Create the name of the service account to use
*/}}
@ -76,10 +123,13 @@ so fall back to "default" (or an explicit override) to avoid a cyclic dependency
{{- end }}
{{/*
Get redis service name
Get redis service name.
The bundled Redis subchart only serves sentinel in "replication" architecture
(it rejects standalone + sentinel outright), and in that mode the sentinel
Service is named "<release>-redis", not "<release>-redis-master".
*/}}
{{- define "litellm.redis.serviceName" -}}
{{- if and (eq .Values.redis.architecture "standalone") .Values.redis.sentinel.enabled -}}
{{- if .Values.redis.sentinel.enabled -}}
{{- printf "%s-%s" .Release.Name (default "redis" .Values.redis.nameOverride | trunc 63 | trimSuffix "-") -}}
{{- else -}}
{{- printf "%s-%s-master" .Release.Name (default "redis" .Values.redis.nameOverride | trunc 63 | trimSuffix "-") -}}

View file

@ -1,9 +1,22 @@
{{- if .Values.proxyConfigMap.create }}
{{- $config := deepCopy .Values.proxy_config }}
{{- if and .Values.redis.enabled (dig "coordination" "enabled" true .Values.redis) }}
{{- $generalSettings := (get $config "general_settings") | default dict }}
{{- if not (hasKey $generalSettings "coordination_redis") }}
{{- $coordinationRedis := dict "host" "os.environ/REDIS_HOST" "port" "os.environ/REDIS_PORT" "password" "os.environ/REDIS_PASSWORD" }}
{{- if .Values.redis.sentinel.enabled }}
{{- $sentinelNode := list (include "litellm.redis.serviceName" .) (include "litellm.redis.port" . | int) }}
{{- $coordinationRedis = dict "sentinel_nodes" (list $sentinelNode) "service_name" (default "mymaster" .Values.redis.sentinel.masterSet) "password" "os.environ/REDIS_PASSWORD" }}
{{- end }}
{{- $_ := set $generalSettings "coordination_redis" $coordinationRedis }}
{{- $_ := set $config "general_settings" $generalSettings }}
{{- end }}
{{- end }}
apiVersion: v1
kind: ConfigMap
metadata:
name: {{ include "litellm.fullname" . }}-config
data:
config.yaml: |
{{ .Values.proxy_config | toYaml | indent 6 }}
{{ $config | toYaml | indent 6 }}
{{- end }}

View file

@ -142,6 +142,9 @@ spec:
{{- with .Values.extraEnvVars }}
{{- toYaml . | nindent 12 }}
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsEnv" . | nindent 12 }}
{{- end }}
{{- if .Values.migrationJob.enabled }}
# Schema updates are owned by the dedicated migrations Job; skip
# the proxy's startup `prisma db push` so N replicas don't race
@ -220,6 +223,9 @@ spec:
- name: npm
mountPath: /.npm
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }}
{{- end }}
{{- with .Values.volumeMounts }}
{{- toYaml . | nindent 12 }}
{{- end }}
@ -252,6 +258,9 @@ spec:
items:
- key: {{ .Values.proxyConfigMap.key | default "config.yaml" }}
path: "config.yaml"
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumes" . | nindent 8 }}
{{- end }}
{{- with .Values.volumes }}
{{- toYaml . | nindent 8 }}
{{- end }}

View file

@ -0,0 +1,297 @@
suite: test billingMetrics wiring on the proxy deployment
templates:
- deployment.yaml
- configmap-litellm.yaml
- migrations-job.yaml
tests:
- it: is off by default, adding no env, volume, or mount
template: deployment.yaml
asserts:
- notContains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: litellm-billing-metrics-mtls
- notContains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls
mountPath: /etc/litellm/billing-mtls
readOnly: true
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- it: renders the endpoint and the mounted cert paths when enabled
template: deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CLIENT_CERT
value: /etc/litellm/billing-mtls/tls.crt
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CLIENT_KEY
value: /etc/litellm/billing-mtls/tls.key
# The conventional Secret name is the default, so enabling the block is enough.
- it: mounts the default cert secret read-only alongside the config volume
template: deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: litellm-billing-metrics-mtls
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls
mountPath: /etc/litellm/billing-mtls
readOnly: true
- it: honours a secretName override
template: deployment.yaml
set:
billingMetrics:
enabled: true
secretName: my-billing-mtls
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: my-billing-mtls
- notContains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: litellm-billing-metrics-mtls
- it: honours an endpoint override
template: deployment.yaml
set:
billingMetrics:
enabled: true
endpoint: https://collector.internal:4318
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://collector.internal:4318
# The production collector presents a public web-PKI certificate, so the CA
# override must stay absent unless a private collector is configured.
- it: omits the CA env, volume, and mount when no caSecretName is set
template: deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- notContains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls-ca
secret:
secretName: billing-ca
- notContains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls-ca
mountPath: /etc/litellm/billing-mtls-ca
readOnly: true
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CA_CERT
value: /etc/litellm/billing-mtls-ca/ca.crt
- it: mounts the CA secret when caSecretName is set
template: deployment.yaml
set:
billingMetrics:
enabled: true
caSecretName: billing-ca
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CA_CERT
value: /etc/litellm/billing-mtls-ca/ca.crt
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls-ca
secret:
secretName: billing-ca
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls-ca
mountPath: /etc/litellm/billing-mtls-ca
readOnly: true
- it: passes the export interval through only when set
template: deployment.yaml
set:
billingMetrics:
enabled: true
exportIntervalMs: 5000
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS
value: "5000"
- it: omits the export interval when unset
template: deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS
value: "60000"
# Kubernetes resolves duplicate env names last-wins, so the chart-owned billing
# entries must render after .Values.envVars or a user could silently redirect
# the metering export. The three billing entries are the last ones emitted here
# (migrationJob, which appends DISABLE_SCHEMA_UPDATE, is off for this case).
- it: renders the billing endpoint after envVars so it cannot be shadowed
template: deployment.yaml
set:
migrationJob:
enabled: false
billingMetrics:
enabled: true
envVars:
LITELLM_BILLING_METRICS_ENDPOINT: https://shadowed.example
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://shadowed.example
- equal:
path: spec.template.spec.containers[0].env[-3]
value:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- equal:
path: spec.template.spec.containers[0].env[-2].name
value: LITELLM_BILLING_METRICS_CLIENT_CERT
- equal:
path: spec.template.spec.containers[0].env[-1].name
value: LITELLM_BILLING_METRICS_CLIENT_KEY
- it: keeps user-supplied volumes and mounts alongside the billing secret
template: deployment.yaml
set:
billingMetrics:
enabled: true
volumes:
- name: custom-callbacks
configMap:
name: my-callbacks
volumeMounts:
- name: custom-callbacks
mountPath: /app/callbacks
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: custom-callbacks
configMap:
name: my-callbacks
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: litellm-billing-metrics-mtls
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: custom-callbacks
mountPath: /app/callbacks
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls
mountPath: /etc/litellm/billing-mtls
readOnly: true
- it: still mounts the proxy config when enabled
template: deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: litellm-config
mountPath: /etc/litellm/config.yaml
subPath: config.yaml
# Only the proxy serves billable traffic. The migrations Job must never mount
# the client certificate, and it renders its own env and volumes, so nothing
# stops a future edit from wiring the billing include into it by mistake.
- it: does not touch the migrations job when enabled
template: migrations-job.yaml
set:
billingMetrics:
enabled: true
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- notExists:
path: spec.template.spec.containers[0].volumeMounts
- notExists:
path: spec.template.spec.volumes
- it: fails loudly when enabled with an emptied secretName
template: deployment.yaml
set:
billingMetrics:
enabled: true
secretName: ""
asserts:
- failedTemplate:
errorMessage: billingMetrics.secretName is required when billingMetrics.enabled is true (an existing Secret with tls.crt and tls.key)
- it: fails loudly when enabled without an endpoint
template: deployment.yaml
set:
billingMetrics:
enabled: true
endpoint: ""
asserts:
- failedTemplate:
errorMessage: billingMetrics.endpoint is required when billingMetrics.enabled is true

View file

@ -0,0 +1,143 @@
suite: test coordination redis
templates:
- configmap-litellm.yaml
- deployment.yaml
tests:
- it: should not render coordination_redis when redis is disabled
template: configmap-litellm.yaml
set:
redis.enabled: false
asserts:
- notMatchRegex:
path: data["config.yaml"]
pattern: coordination_redis
- it: should not emit redis env vars when redis is disabled
template: deployment.yaml
set:
redis.enabled: false
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: RELEASE-NAME-redis-master
any: true
- it: should render coordination_redis pointing at the bundled redis when enabled
template: configmap-litellm.yaml
set:
redis.enabled: true
asserts:
- matchRegex:
path: data["config.yaml"]
pattern: "coordination_redis:\n host: os.environ/REDIS_HOST\n password: os.environ/REDIS_PASSWORD\n port: os.environ/REDIS_PORT\n"
- matchRegex:
path: data["config.yaml"]
pattern: "master_key: os.environ/PROXY_MASTER_KEY"
- it: should emit redis env vars backing the coordination_redis os.environ refs
template: deployment.yaml
set:
redis.enabled: true
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: RELEASE-NAME-redis-master
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PORT
value: "6379"
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PASSWORD
valueFrom:
secretKeyRef:
name: RELEASE-NAME-redis
key: redis-password
- it: should not render coordination_redis when coordination is opted out
template: configmap-litellm.yaml
set:
redis.enabled: true
redis.coordination.enabled: false
asserts:
- notMatchRegex:
path: data["config.yaml"]
pattern: coordination_redis
- it: should keep emitting redis env vars when coordination is opted out
template: deployment.yaml
set:
redis.enabled: true
redis.coordination.enabled: false
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: RELEASE-NAME-redis-master
- it: should not clobber a user supplied coordination_redis block
template: configmap-litellm.yaml
set:
redis.enabled: true
proxy_config.general_settings.coordination_redis:
url: os.environ/COORDINATION_REDIS_URL
asserts:
- matchRegex:
path: data["config.yaml"]
pattern: "coordination_redis:\n url: os.environ/COORDINATION_REDIS_URL\n"
- notMatchRegex:
path: data["config.yaml"]
pattern: "host: os.environ/REDIS_HOST"
- it: should render sentinel_nodes and service_name in sentinel mode
template: configmap-litellm.yaml
set:
redis.enabled: true
redis.architecture: replication
redis.sentinel.enabled: true
asserts:
# The sentinel Service the redis subchart renders is "<release>-redis", and a
# plain client cannot speak the sentinel protocol, so host/port must not appear
- matchRegex:
path: data["config.yaml"]
pattern: "coordination_redis:\n password: os.environ/REDIS_PASSWORD\n sentinel_nodes:\n - - RELEASE-NAME-redis\n - 26379\n service_name: mymaster\n"
- notMatchRegex:
path: data["config.yaml"]
pattern: "host: os.environ/REDIS_HOST"
- it: should carry a custom sentinel masterSet into service_name
template: configmap-litellm.yaml
set:
redis.enabled: true
redis.architecture: replication
redis.sentinel.enabled: true
redis.sentinel.masterSet: litellm-master
asserts:
- matchRegex:
path: data["config.yaml"]
pattern: "service_name: litellm-master"
- it: should point REDIS_HOST at the sentinel service in sentinel mode
template: deployment.yaml
set:
redis.enabled: true
redis.architecture: replication
redis.sentinel.enabled: true
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: RELEASE-NAME-redis
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PORT
value: "26379"

View file

@ -139,6 +139,20 @@ masterkeySecretName: ""
# if set, use this secret key for the master key; otherwise, use the default key
masterkeySecretKey: ""
# Optional: enterprise billable-request metering. When enabled, the proxy counts
# successful requests to inference, MCP, and A2A endpoints and pushes them to
# LiteLLM's collector over mutual TLS. Requires an enterprise license.
# The client certificate identifies the deployment, so it is mounted read-only
# from an existing Secret and never passed through the environment.
billingMetrics:
enabled: false
endpoint: https://telemetry.litellm.ai # collector to push the counter to
secretName: litellm-billing-metrics-mtls # existing Secret holding tls.crt and tls.key
# Only for private or test collectors whose server certificate is not on the
# public web PKI. The production collector needs no CA override.
caSecretName: "" # existing Secret holding ca.crt
exportIntervalMs: "" # push cadence; the proxy defaults to 60000
proxyConfigMap:
# when true, creates a new configmap
create: true
@ -331,12 +345,28 @@ postgresql:
# secretKeys:
# userPasswordKey: password
# requires cache: true in config file
# either enable this or pass a secret for REDIS_HOST, REDIS_PORT, REDIS_PASSWORD or REDIS_URL
# with cache: true to use existing redis instance
# Redis is the proxy's coordination store: cross-pod tpm/rpm rate limits, spend
# tracking, and the pod lock manager. Enabling this deploys the bundled Redis
# subchart, wires REDIS_HOST / REDIS_PORT / REDIS_PASSWORD into the proxy, and
# renders a `general_settings.coordination_redis` block into the proxy config.
#
# To point at an existing Redis instead, leave `enabled: false` and pass a
# secret for REDIS_HOST, REDIS_PORT, REDIS_PASSWORD or REDIS_URL; the proxy
# falls back to those env vars for coordination. Set `cache: true` in the proxy
# config only if you also want LLM response caching, which is independent of
# coordination
#
# When `redis.sentinel.enabled` is set, the coordination block is rendered with
# `sentinel_nodes` and `service_name` (from `redis.sentinel.masterSet`) instead
# of host/port, because a plain Redis client cannot talk to the sentinel port
redis:
enabled: false
architecture: standalone
coordination:
# Set to false to keep the bundled Redis for response caching only and leave
# `general_settings.coordination_redis` out of the rendered config. A
# `coordination_redis` block you define yourself in `proxy_config` always wins
enabled: true
# Prisma migration job settings
migrationJob:

View file

@ -46,4 +46,9 @@ Reminders:
- gateway.config.proxy_config (rendered into a ConfigMap and mounted at
/app/config/config.yaml; gateway reads it via
CONFIG_FILE_PATH)
- {component}.pdb.{enabled,minAvailable,maxUnavailable} (per-component PodDisruptionBudget; disabled by
default — with hpa.minReplicas of 1, minAvailable: 1
would block node drains)
- {component}.topologySpreadConstraints (standard k8s list, e.g. spread replicas across
topology.kubernetes.io/zone)
- Enable ingress.enabled=true to dispatch / → ui, gateway data-plane prefixes → gateway, and the catch-all → backend.

View file

@ -34,6 +34,57 @@ app.kubernetes.io/managed-by: {{ .Release.Service }}
helm.sh/chart: {{ printf "%s-%s" .Chart.Name .Chart.Version | replace "+" "_" }}
{{- end -}}
{{/*
Enterprise billable-request metering. Wired into gateway and backend, not the
migrations job. The gateway serves nearly all billable traffic, but the backend
keeps the named-server MCP transport (/{mcp_server_name}/mcp), which writes a
SpendLogs row, so metering only the gateway would silently drop that traffic.
The client certificate identifies the deployment to LiteLLM's collector, so it is
mounted read-only from an existing Secret rather than passed through the
environment.
*/}}
{{- define "litellm.billingMetrics.certDir" -}}/etc/litellm/billing-mtls{{- end -}}
{{- define "litellm.billingMetrics.caDir" -}}/etc/litellm/billing-mtls-ca{{- end -}}
{{- define "litellm.billingMetricsEnv" -}}
- name: LITELLM_BILLING_METRICS_ENDPOINT
value: {{ required "billingMetrics.endpoint is required when billingMetrics.enabled is true" .Values.billingMetrics.endpoint | quote }}
- name: LITELLM_BILLING_METRICS_CLIENT_CERT
value: {{ printf "%s/tls.crt" (include "litellm.billingMetrics.certDir" .) | quote }}
- name: LITELLM_BILLING_METRICS_CLIENT_KEY
value: {{ printf "%s/tls.key" (include "litellm.billingMetrics.certDir" .) | quote }}
{{- if .Values.billingMetrics.caSecretName }}
- name: LITELLM_BILLING_METRICS_CA_CERT
value: {{ printf "%s/ca.crt" (include "litellm.billingMetrics.caDir" .) | quote }}
{{- end }}
{{- with .Values.billingMetrics.exportIntervalMs }}
- name: LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS
value: {{ . | quote }}
{{- end }}
{{- end -}}
{{- define "litellm.billingMetricsVolumes" -}}
- name: billing-metrics-mtls
secret:
secretName: {{ required "billingMetrics.secretName is required when billingMetrics.enabled is true (an existing Secret with tls.crt and tls.key)" .Values.billingMetrics.secretName }}
{{- if .Values.billingMetrics.caSecretName }}
- name: billing-metrics-mtls-ca
secret:
secretName: {{ .Values.billingMetrics.caSecretName }}
{{- end }}
{{- end -}}
{{- define "litellm.billingMetricsVolumeMounts" -}}
- name: billing-metrics-mtls
mountPath: {{ include "litellm.billingMetrics.certDir" . }}
readOnly: true
{{- if .Values.billingMetrics.caSecretName }}
- name: billing-metrics-mtls-ca
mountPath: {{ include "litellm.billingMetrics.caDir" . }}
readOnly: true
{{- end }}
{{- end -}}
{{/*
Per-component selector labels — used in both Service selectors and Deployment matchLabels.
*/}}
@ -213,6 +264,10 @@ harmless no-op for the Job and authoritative for the app pods.
*/}}
- name: DISABLE_SCHEMA_UPDATE
value: "true"
{{/* These feed the proxy's coordination Redis (cross-pod rate limits, spend
tracking, pod lock manager) via its REDIS_* env fallback. An explicit
`general_settings.coordination_redis` block in proxy_config takes
precedence over anything emitted here. */}}
{{- if $root.Values.redis.host }}
- name: REDIS_HOST
value: {{ $root.Values.redis.host | quote }}
@ -226,10 +281,11 @@ harmless no-op for the Job and authoritative for the app pods.
key: {{ $root.Values.redis.passwordSecret.passwordKey | default "password" }}
{{- end }}
{{- if $root.Values.redis.cluster }}
{{/* The proxy's Cache() reads REDIS_CLUSTER_NODES as JSON and constructs a
RedisClusterCache when it's set (litellm/caching/caching.py:169-192).
We seed with the single configured endpoint — the cluster client
discovers the remaining nodes from CLUSTER SLOTS at startup. */}}
{{/* The proxy falls back to REDIS_CLUSTER_NODES (JSON) to build a cluster-mode
coordination client when `general_settings.coordination_redis` is absent
and no plain-Redis response cache is configured. We seed with the single
configured endpoint; the cluster client discovers the remaining nodes from
CLUSTER SLOTS at startup. */}}
- name: REDIS_CLUSTER_NODES
value: {{ printf "[{\"host\":%q,\"port\":%v}]" $root.Values.redis.host (int $root.Values.redis.port) | quote }}
{{- end }}
@ -239,6 +295,52 @@ harmless no-op for the Job and authoritative for the app pods.
{{- end }}
{{- end -}}
{{/*
PodDisruptionBudget shared by gateway, backend, and ui.
Invoke with a dict:
(dict "root" $ "component" .Values.gateway "componentName" "gateway"
"fullname" (include "litellm.gateway.fullname" .)
"selectorLabels" (include "litellm.gateway.selectorLabels" .))
Renders nothing unless both the component and its `pdb.enabled` are on.
Only one of minAvailable / maxUnavailable should be set; if both are,
minAvailable wins. If neither is set, falls back to `maxUnavailable: 1` so
an enabled-but-unconfigured PDB still permits node drains.
"Set" means non-nil and non-empty-string, so an explicit 0 (e.g.
`maxUnavailable: 0` to forbid all voluntary disruptions) is honored rather
than silently replaced by the fallback.
*/}}
{{- define "litellm.pdb" -}}
{{- $root := .root -}}
{{- $component := .component -}}
{{- $min := $component.pdb.minAvailable -}}
{{- $max := $component.pdb.maxUnavailable -}}
{{- $minSet := not (or (kindIs "invalid" $min) (eq (printf "%v" $min) "")) -}}
{{- $maxSet := not (or (kindIs "invalid" $max) (eq (printf "%v" $max) "")) -}}
{{- if and $component.enabled $component.pdb $component.pdb.enabled }}
apiVersion: policy/v1
kind: PodDisruptionBudget
metadata:
name: {{ .fullname }}
labels:
{{- include "litellm.commonLabels" $root | nindent 4 }}
app.kubernetes.io/component: {{ .componentName }}
spec:
selector:
matchLabels:
{{- .selectorLabels | nindent 6 }}
{{- if $minSet }}
minAvailable: {{ $min }}
{{- else if $maxSet }}
maxUnavailable: {{ $max }}
{{- else }}
maxUnavailable: 1
{{- end }}
{{- end }}
{{- end -}}
{{/*
Renders `envFrom:` block for a component's `envConfigMaps` / `envSecrets`
lists. Each entry is a resource name; the chart wires the whole ConfigMap /

View file

@ -44,14 +44,20 @@ spec:
- name: CONFIG_FILE_PATH
value: /app/config/config.yaml
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsEnv" . | nindent 12 }}
{{- end }}
{{- include "litellm.envFrom" .Values.backend | nindent 10 }}
{{- if or .Values.gateway.config.create .Values.backend.volumeMounts }}
{{- if or .Values.gateway.config.create .Values.backend.volumeMounts .Values.billingMetrics.enabled }}
volumeMounts:
{{- if .Values.gateway.config.create }}
- name: gateway-config
mountPath: /app/config/config.yaml
subPath: config.yaml
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }}
{{- end }}
{{- with .Values.backend.volumeMounts }}
{{- toYaml . | nindent 12 }}
{{- end }}
@ -66,13 +72,16 @@ spec:
{{- end }}
resources:
{{- toYaml .Values.backend.resources | nindent 12 }}
{{- if or .Values.gateway.config.create .Values.backend.volumes }}
{{- if or .Values.gateway.config.create .Values.backend.volumes .Values.billingMetrics.enabled }}
volumes:
{{- if .Values.gateway.config.create }}
- name: gateway-config
configMap:
name: {{ include "litellm.gateway.fullname" . }}-config
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumes" . | nindent 8 }}
{{- end }}
{{- with .Values.backend.volumes }}
{{- toYaml . | nindent 8 }}
{{- end }}
@ -89,4 +98,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.backend.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.backend
"componentName" "backend"
"fullname" (include "litellm.backend.fullname" .)
"selectorLabels" (include "litellm.backend.selectorLabels" .)) }}

View file

@ -46,14 +46,20 @@ spec:
- name: NUM_WORKERS
value: {{ .Values.gateway.numWorkers | quote }}
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsEnv" . | nindent 12 }}
{{- end }}
{{- include "litellm.envFrom" .Values.gateway | nindent 10 }}
{{- if or .Values.gateway.config.create .Values.gateway.volumeMounts }}
{{- if or .Values.gateway.config.create .Values.gateway.volumeMounts .Values.billingMetrics.enabled }}
volumeMounts:
{{- if .Values.gateway.config.create }}
- name: gateway-config
mountPath: /app/config/config.yaml
subPath: config.yaml
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }}
{{- end }}
{{- with .Values.gateway.volumeMounts }}
{{- toYaml . | nindent 12 }}
{{- end }}
@ -68,13 +74,16 @@ spec:
{{- end }}
resources:
{{- toYaml .Values.gateway.resources | nindent 12 }}
{{- if or .Values.gateway.config.create .Values.gateway.volumes }}
{{- if or .Values.gateway.config.create .Values.gateway.volumes .Values.billingMetrics.enabled }}
volumes:
{{- if .Values.gateway.config.create }}
- name: gateway-config
configMap:
name: {{ include "litellm.gateway.fullname" . }}-config
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumes" . | nindent 8 }}
{{- end }}
{{- with .Values.gateway.volumes }}
{{- toYaml . | nindent 8 }}
{{- end }}
@ -91,4 +100,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.gateway.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.gateway
"componentName" "gateway"
"fullname" (include "litellm.gateway.fullname" .)
"selectorLabels" (include "litellm.gateway.selectorLabels" .)) }}

View file

@ -76,4 +76,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.ui.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.ui
"componentName" "ui"
"fullname" (include "litellm.ui.fullname" .)
"selectorLabels" (include "litellm.ui.selectorLabels" .)) }}

View file

@ -0,0 +1,249 @@
suite: test billingMetrics wiring on gateway and backend
templates:
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
- migrations-job.yaml
values:
- ./values/required.yaml
tests:
- it: is off by default, adding no env, volume, or mount
template: gateway/deployment.yaml
asserts:
- notContains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: billing-mtls
- equal:
path: spec.template.spec.containers[0].volumeMounts
value:
- name: gateway-config
mountPath: /app/config/config.yaml
subPath: config.yaml
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- it: renders the endpoint and the mounted cert paths when enabled
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CLIENT_CERT
value: /etc/litellm/billing-mtls/tls.crt
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CLIENT_KEY
value: /etc/litellm/billing-mtls/tls.key
- it: mounts the cert secret read-only alongside the config volume
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: billing-mtls
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls
mountPath: /etc/litellm/billing-mtls
readOnly: true
# The production collector presents a public web-PKI certificate, so the CA
# override must stay absent unless a private collector is configured.
- it: omits the CA env, volume, and mount when no caSecretName is set
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
asserts:
- notContains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls-ca
secret:
secretName: billing-ca
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CA_CERT
value: /etc/litellm/billing-mtls-ca/ca.crt
- it: mounts the CA secret when caSecretName is set
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
caSecretName: billing-ca
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_CA_CERT
value: /etc/litellm/billing-mtls-ca/ca.crt
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls-ca
secret:
secretName: billing-ca
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls-ca
mountPath: /etc/litellm/billing-mtls-ca
readOnly: true
- it: passes the export interval through only when set
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
exportIntervalMs: 5000
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS
value: "5000"
- it: keeps user-supplied gateway volumes alongside the billing secret
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
gateway.volumes:
- name: custom-callbacks
configMap:
name: my-callbacks
gateway.volumeMounts:
- name: custom-callbacks
mountPath: /app/callbacks
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: custom-callbacks
configMap:
name: my-callbacks
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: billing-mtls
# The backend keeps the named-server MCP transport (/{mcp_server_name}/mcp),
# which writes a SpendLogs row, so it must meter too or that traffic is lost.
- it: meters the backend as well, since it serves the MCP transport
template: backend/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: billing-metrics-mtls
mountPath: /etc/litellm/billing-mtls
readOnly: true
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: billing-mtls
- it: leaves the backend alone when metering is off
template: backend/deployment.yaml
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
# The migrations job runs prisma and serves no traffic; it must never receive
# the client key.
- it: never mounts the billing cert on the migrations job
template: migrations-job.yaml
set:
billingMetrics:
enabled: true
secretName: billing-mtls
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_BILLING_METRICS_ENDPOINT
value: https://telemetry.litellm.ai
- isNull:
path: spec.template.spec.volumes
# The conventional Secret name is the default, so enabling metering needs no
# secretName at all; the guard below only fires on an explicitly blanked one.
- it: uses the conventional secret name by default
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
asserts:
- contains:
path: spec.template.spec.volumes
content:
name: billing-metrics-mtls
secret:
secretName: litellm-billing-metrics-mtls
- it: fails loudly when the secretName is explicitly blanked
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
secretName: ""
asserts:
- failedTemplate:
errorMessage: billingMetrics.secretName is required when billingMetrics.enabled is true (an existing Secret with tls.crt and tls.key)
- it: fails loudly when enabled without an endpoint
template: gateway/deployment.yaml
set:
billingMetrics:
enabled: true
endpoint: ""
secretName: billing-mtls
asserts:
- failedTemplate:
errorMessage: billingMetrics.endpoint is required when billingMetrics.enabled is true

View file

@ -0,0 +1,188 @@
suite: test pod disruption budgets and topology spread constraints
templates:
- gateway/poddisruptionbudget.yaml
- backend/poddisruptionbudget.yaml
- ui/poddisruptionbudget.yaml
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
- ui/deployment.yaml
values:
- ./values/required.yaml
tests:
- it: renders no PDB by default
templates:
- gateway/poddisruptionbudget.yaml
- backend/poddisruptionbudget.yaml
- ui/poddisruptionbudget.yaml
asserts:
- hasDocuments:
count: 0
- it: gateway PDB uses minAvailable and matches the gateway selector labels
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 1
asserts:
- isKind:
of: PodDisruptionBudget
- equal:
path: apiVersion
value: policy/v1
- equal:
path: metadata.name
value: RELEASE-NAME-litellm-gateway
- equal:
path: spec.minAvailable
value: 1
- notExists:
path: spec.maxUnavailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: gateway
- it: backend PDB uses maxUnavailable when minAvailable is unset
template: backend/poddisruptionbudget.yaml
set:
backend.pdb.enabled: true
backend.pdb.maxUnavailable: 25%
asserts:
- equal:
path: spec.maxUnavailable
value: 25%
- notExists:
path: spec.minAvailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: backend
- it: minAvailable wins when both minAvailable and maxUnavailable are set
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 2
gateway.pdb.maxUnavailable: 1
asserts:
- equal:
path: spec.minAvailable
value: 2
- notExists:
path: spec.maxUnavailable
- it: an explicit maxUnavailable 0 is honored instead of the fallback
template: backend/poddisruptionbudget.yaml
set:
backend.pdb.enabled: true
backend.pdb.maxUnavailable: 0
asserts:
- equal:
path: spec.maxUnavailable
value: 0
- notExists:
path: spec.minAvailable
- it: an explicit minAvailable 0 is honored and beats a set maxUnavailable
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 0
gateway.pdb.maxUnavailable: 1
asserts:
- equal:
path: spec.minAvailable
value: 0
- notExists:
path: spec.maxUnavailable
- it: enabled PDB with neither knob set falls back to maxUnavailable 1
template: ui/poddisruptionbudget.yaml
set:
ui.pdb.enabled: true
asserts:
- equal:
path: spec.maxUnavailable
value: 1
- notExists:
path: spec.minAvailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: ui
- it: renders no PDB for a disabled component even when its pdb is enabled
template: gateway/poddisruptionbudget.yaml
set:
gateway.enabled: false
gateway.pdb.enabled: true
asserts:
- hasDocuments:
count: 0
- it: deployments omit topologySpreadConstraints by default
templates:
- gateway/deployment.yaml
- backend/deployment.yaml
- ui/deployment.yaml
asserts:
- notExists:
path: spec.template.spec.topologySpreadConstraints
- it: gateway deployment renders configured topologySpreadConstraints
template: gateway/deployment.yaml
set:
gateway.topologySpreadConstraints:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
labelSelector:
matchLabels:
app.kubernetes.io/component: gateway
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints
value:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
labelSelector:
matchLabels:
app.kubernetes.io/component: gateway
- it: backend deployment renders configured topologySpreadConstraints
template: backend/deployment.yaml
set:
backend.topologySpreadConstraints:
- maxSkew: 1
topologyKey: kubernetes.io/hostname
whenUnsatisfiable: DoNotSchedule
labelSelector:
matchLabels:
app.kubernetes.io/component: backend
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints[0].topologyKey
value: kubernetes.io/hostname
- equal:
path: spec.template.spec.topologySpreadConstraints[0].whenUnsatisfiable
value: DoNotSchedule
- it: ui deployment renders configured topologySpreadConstraints
template: ui/deployment.yaml
set:
ui.topologySpreadConstraints:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints[0].topologyKey
value: topology.kubernetes.io/zone

View file

@ -0,0 +1,109 @@
suite: test redis coordination env vars
templates:
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
values:
- ./values/required.yaml
tests:
- it: gateway omits redis env vars when no host is configured
template: gateway/deployment.yaml
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: redis.example.com
any: true
- notContains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_CLUSTER_NODES
any: true
- it: gateway emits host, port and password when redis is configured
template: gateway/deployment.yaml
set:
redis.host: redis.example.com
redis.port: 6380
redis.passwordSecret.name: redis-secret
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: redis.example.com
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PORT
value: "6380"
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PASSWORD
valueFrom:
secretKeyRef:
name: redis-secret
key: password
- it: backend emits the same redis env vars so both pods coordinate on one redis
template: backend/deployment.yaml
set:
redis.host: redis.example.com
redis.passwordSecret.name: redis-secret
redis.passwordSecret.passwordKey: redis-password
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: redis.example.com
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PASSWORD
valueFrom:
secretKeyRef:
name: redis-secret
key: redis-password
- it: gateway omits REDIS_PASSWORD for an auth-less redis
template: gateway/deployment.yaml
set:
redis.host: redis.example.com
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_PASSWORD
any: true
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_HOST
value: redis.example.com
- it: gateway seeds REDIS_CLUSTER_NODES from host and port in cluster mode
template: gateway/deployment.yaml
set:
redis.host: redis.example.com
redis.port: 6380
redis.cluster: true
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_CLUSTER_NODES
value: '[{"host":"redis.example.com","port":6380}]'
- it: gateway omits REDIS_CLUSTER_NODES when cluster mode is off
template: gateway/deployment.yaml
set:
redis.host: redis.example.com
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: REDIS_CLUSTER_NODES
any: true

View file

@ -73,6 +73,25 @@ masterKey:
secretName: litellm-master-key-secret # name of a Secret containing the master key
secretKey: master-key
# Optional: enterprise billable-request metering. When enabled, the gateway and
# backend count successful requests to inference, MCP, and A2A endpoints and push
# them to LiteLLM's collector over mutual TLS. Both components serve billable
# routes: the backend keeps the named-server MCP transport. Requires an
# enterprise license. The client certificate identifies the deployment, so it is
# mounted read-only from an existing Secret and never passed through the env.
billingMetrics:
enabled: false
endpoint: https://telemetry.litellm.ai # collector to push the counter to
# An existing Secret holding the client certificate under tls.crt and its key
# under tls.key, usually created from the onboarding artifact. The default is
# the conventional name, so the common path is to create that Secret and set
# enabled: true. Override only if yours is named differently.
secretName: litellm-billing-metrics-mtls
# Only for private or test collectors whose server certificate is not on the
# public web PKI. The production collector needs no CA override.
caSecretName: "" # existing Secret holding ca.crt
exportIntervalMs: "" # push cadence; the proxy defaults to 60000
# External Postgres connection.
database:
writer:
@ -100,7 +119,18 @@ database:
usernameKey: username
passwordKey: password
# Optional Redis (caching, rate limiting). Leave host empty to disable.
# Optional Redis. Leave host empty to disable.
#
# This is the proxy's coordination store: cross-pod tpm/rpm rate limits, spend
# tracking, and the pod lock manager. The chart emits REDIS_HOST / REDIS_PORT /
# REDIS_PASSWORD, which the proxy picks up through its coordination Redis env
# fallback. Response caching is separate and off unless you enable it in
# `proxy_config.litellm_settings.cache`.
#
# For full control, define `general_settings.coordination_redis` in
# `proxy_config` (host/port/password/username/url/ssl/startup_nodes/
# sentinel_nodes/sentinel_password/service_name, each accepting os.environ/VAR
# refs). An explicit block overrides these env vars.
#
# Set `cluster: true` for Redis Cluster mode (e.g. AWS ElastiCache Cluster,
# self-hosted Redis Cluster). The chart emits REDIS_CLUSTER_NODES from
@ -160,10 +190,28 @@ gateway:
maxReplicas: 10
targetCPUUtilizationPercentage: 70
targetMemoryUtilizationPercentage: 80
# PodDisruptionBudget for the gateway pods. Set exactly one of
# `minAvailable` / `maxUnavailable` (minAvailable wins if both are set;
# enabling without either falls back to `maxUnavailable: 1`). Disabled by
# default: with the default hpa.minReplicas of 1, a `minAvailable: 1` PDB
# would block node drains entirely.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Standard k8s topologySpreadConstraints for the gateway pods, e.g. to
# spread replicas across zones:
# - maxSkew: 1
# topologyKey: topology.kubernetes.io/zone
# whenUnsatisfiable: ScheduleAnyway
# labelSelector:
# matchLabels:
# app.kubernetes.io/component: gateway
topologySpreadConstraints: []
# ---------- backend (UI / management API) ----------
backend:
@ -203,10 +251,17 @@ backend:
minReplicas: 1
maxReplicas: 4
targetCPUUtilizationPercentage: 70
# Same shape as gateway.pdb.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Same shape as gateway.topologySpreadConstraints.
topologySpreadConstraints: []
# ---------- ui (Next.js static dashboard) ----------
ui:
@ -249,7 +304,14 @@ ui:
minReplicas: 1
maxReplicas: 3
targetCPUUtilizationPercentage: 80
# Same shape as gateway.pdb.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Same shape as gateway.topologySpreadConstraints.
topologySpreadConstraints: []

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "dcr_bridge" BOOLEAN;

View file

@ -0,0 +1,6 @@
-- AlterTable
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN "key_type" TEXT;
-- AlterTable
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "key_type" TEXT;

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "issuer" TEXT;

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
command String?
args String[] @default([])
env Json? @default("{}")
issuer String?
authorization_url String?
token_url String?
registration_url String?
@ -339,6 +340,7 @@ model LiteLLM_MCPServerTable {
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
oauth_passthrough Boolean @default(false)
dcr_bridge Boolean?
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?
@ -421,6 +423,7 @@ model LiteLLM_VerificationToken {
budget_reset_at DateTime?
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
key_type String?
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")
@ -515,6 +518,7 @@ model LiteLLM_DeletedVerificationToken {
budget_reset_at DateTime?
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
key_type String?
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.75"
version = "0.4.78"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.75"
version = "0.4.78"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -325,8 +325,19 @@ def _get_redis_client_logic(**env_overrides):
value = get_secret(v) # type: ignore
env_overrides[k] = value
environment_kwargs = _redis_kwargs_from_environment()
# An explicitly configured connection target outranks REDIS_URL from the
# environment. Without this, the url branch below strips the caller's
# host/port/password and silently connects to whatever REDIS_URL names.
caller_named_a_target = any(
env_overrides.get(key) is not None for key in ("host", "startup_nodes", "sentinel_nodes")
)
if caller_named_a_target and env_overrides.get("url") is None:
environment_kwargs.pop("url", None)
redis_kwargs = {
**_redis_kwargs_from_environment(),
**environment_kwargs,
**env_overrides,
}
@ -677,11 +688,8 @@ def get_redis_connection_pool(
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
connection_class = async_redis.Connection
if "ssl" in redis_kwargs:
connection_class = async_redis.SSLConnection
redis_kwargs.pop("ssl", None)
redis_kwargs["connection_class"] = connection_class
if redis_kwargs.pop("ssl", None):
redis_kwargs["connection_class"] = async_redis.SSLConnection
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)

View file

@ -118,6 +118,7 @@ def _batch_cost_calculator(
total_cost = _get_batch_job_cost_from_file_content(
file_content_dictionary=file_content_dictionary,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
model_info=model_info,
)
verbose_logger.debug("total_cost=%s", total_cost)
@ -363,6 +364,7 @@ def _count_entry_tokens(
def _get_batch_job_cost_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> float:
"""
@ -377,9 +379,15 @@ def _get_batch_job_cost_from_file_content(
for _item in file_content_dictionary:
if _batch_response_was_successful(_item, custom_llm_provider):
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
if model_info is not None or custom_llm_provider == "anthropic":
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
model = _response_body.get("model", "")
# Bedrock batch output lines report a short internal model id
# (e.g. "claude-sonnet-4-6") that is not in the cost map; use the
# deployment model name for pricing when available.
if custom_llm_provider == "bedrock" and model_name:
model = model_name
else:
model = _response_body.get("model") or model_name or ""
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=model,
@ -485,7 +493,7 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov
"""
Get the tokens of a batch job from the response body
"""
if custom_llm_provider == "anthropic":
if custom_llm_provider in ("anthropic", "bedrock"):
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
return AnthropicConfig().calculate_usage(
@ -513,6 +521,8 @@ def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom
"""
if custom_llm_provider == "anthropic":
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("message", None) or {}
if custom_llm_provider == "bedrock":
return batch_job_output_file.get("modelOutput", None) or {}
_response: dict = batch_job_output_file.get("response", None) or {}
_response_body = _response.get("body", None) or {}
return _response_body
@ -523,9 +533,12 @@ def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provi
Check if the batch job response was successful
OpenAI-shaped output rows report ``response.status_code == 200``; Anthropic
message batch results lines report ``result.type == "succeeded"``.
message batch results lines report ``result.type == "succeeded"``; Bedrock
batch output lines report ``modelOutput`` (and no ``error``).
"""
if custom_llm_provider == "anthropic":
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("type") == "succeeded"
if custom_llm_provider == "bedrock":
return batch_job_output_file.get("modelOutput") is not None and batch_job_output_file.get("error") is None
_response: dict = batch_job_output_file.get("response", None) or {}
return _response.get("status_code", None) == 200

View file

@ -85,6 +85,22 @@ class CachingHandlerResponse(BaseModel):
in_memory_cache_obj = InMemoryCache()
def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str, object]:
"""
The caching handler is stored on the Logging object
(``logging_obj._llm_caching_handler``), so keeping ``litellm_logging_obj``
inside ``request_kwargs`` closes a reference cycle
(Logging -> LLMCachingHandler -> kwargs -> Logging) that keeps the full
request payload (messages included) alive until a generational GC pass
instead of being freed by refcount when the request ends. Nothing in the
caching layer reads the logging object from these kwargs; cache-key
generation ignores litellm-internal params.
"""
if "litellm_logging_obj" not in request_kwargs:
return request_kwargs
return {k: v for k, v in request_kwargs.items() if k != "litellm_logging_obj"}
def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
cached_id = cached_result.get("id")
if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"):
@ -118,7 +134,7 @@ class LLMCachingHandler:
self.async_streaming_chunks: List[ModelResponse] = []
self.sync_streaming_chunks: List[ModelResponse] = []
self.request_kwargs = request_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(request_kwargs)
self.preset_cache_key: Optional[str] = None
self.original_function = original_function
self.start_time = start_time
@ -297,7 +313,7 @@ class LLMCachingHandler:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
print_verbose("Checking Sync Cache")
cached_result = litellm.cache.get_cache(**new_kwargs)
if cached_result is not None:
@ -693,7 +709,7 @@ class LLMCachingHandler:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
cached_result: Optional[Any] = None
if call_type == CallTypes.aembedding.value:
if isinstance(new_kwargs["input"], str):

View file

@ -103,6 +103,18 @@ class DualCache(BaseCache):
if default_redis_ttl is not None:
self.default_redis_ttl = default_redis_ttl
def _backfill_kwargs(self, kwargs: "dict[str, object]") -> "dict[str, object]":
"""
Kwargs for writing a Redis read result into the in-memory tier.
Applies ``default_in_memory_ttl`` exactly like the write paths do;
without it, backfilled entries fall to ``InMemoryCache``'s own default
TTL and can outlive the TTL this cache was configured with.
"""
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
return {**kwargs, "ttl": self.default_in_memory_ttl}
return kwargs
def set_cache(self, key, value, local_only: bool = False, **kwargs):
# Update both Redis and in-memory cache
try:
@ -160,7 +172,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
self.in_memory_cache.set_cache(key, redis_result, **kwargs)
self.in_memory_cache.set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -226,7 +238,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
await self.in_memory_cache.async_set_cache(key, redis_result, **kwargs)
await self.in_memory_cache.async_set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -318,7 +330,7 @@ class DualCache(BaseCache):
result[key_to_index[key]] = value
if value is not None and self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return result
except Exception:

View file

@ -2,7 +2,7 @@ import os
import sys
from typing import List, Literal, Optional
from litellm.litellm_core_utils.env_utils import get_env_int
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none
DEFAULT_HEALTH_CHECK_PROMPT = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
AZURE_DEFAULT_RESPONSES_API_VERSION = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
@ -269,9 +269,18 @@ TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 6
MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8)))
MAX_IN_MEMORY_QUEUE_FLUSH_COUNT = int(os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000))
###############################################################################################
MINIMUM_PROMPT_CACHE_TOKEN_COUNT = int(
os.getenv("MINIMUM_PROMPT_CACHE_TOKEN_COUNT", 1024)
) # minimum number of tokens to cache a prompt by Anthropic
# Providers will not cache a prefix below a minimum size. That minimum is per-model, not global:
# Anthropic's ranges from 512 to 4096 depending on the model, and can differ per platform for the
# same model. The real minimum is resolved from `prompt_cache_min_tokens` in the model cost map;
# this value is only the fallback for models the cost map has no entry for, and doubles as a global
# escape hatch when `MINIMUM_PROMPT_CACHE_TOKEN_COUNT` is explicitly set.
MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE: int | None = get_env_int_or_none("MINIMUM_PROMPT_CACHE_TOKEN_COUNT")
DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT = 1024
MINIMUM_PROMPT_CACHE_TOKEN_COUNT = (
MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE
if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None
else DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
)
DEFAULT_TRIM_RATIO = float(
os.getenv("DEFAULT_TRIM_RATIO", 0.75)
) # default ratio of tokens to trim from the end of a prompt
@ -1496,6 +1505,7 @@ MAX_TEAM_LIST_LIMIT = int(os.getenv("MAX_TEAM_LIST_LIMIT", 20))
MAX_POLICY_ESTIMATE_IMPACT_ROWS = int(os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000))
DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD", 0.7))
LENGTH_OF_LITELLM_GENERATED_KEY = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY", 16))
MINIMUM_CUSTOM_KEY_LENGTH = int(os.getenv("MINIMUM_CUSTOM_KEY_LENGTH", 16))
SECRET_MANAGER_REFRESH_INTERVAL = int(os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400))
LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
"default_internal_user_params",

View file

@ -760,7 +760,11 @@ def _select_model_name_for_cost_calc(
if custom_pricing is True:
if router_model_id is not None and router_model_id in litellm.model_cost:
entry = litellm.model_cost[router_model_id]
if entry.get("input_cost_per_token") is not None or entry.get("input_cost_per_second") is not None:
if (
entry.get("input_cost_per_token") is not None
or entry.get("input_cost_per_second") is not None
or entry.get("tiered_pricing") is not None
):
return_model = router_model_id
else:
return_model = model

View file

@ -966,11 +966,15 @@ class BudgetExceededError(Exception):
max_budget: float,
message: Optional[str] = None,
llm_provider: Optional[str] = None,
entity_type: Optional[str] = None,
entity_id: Optional[str] = None,
):
self.current_cost = current_cost
self.max_budget = max_budget
self.status_code = 429
self.llm_provider = llm_provider or ""
self.entity_type = entity_type
self.entity_id = entity_id
# Surface unified rate-limit fields without joining the RateLimitError
# hierarchy so existing `except BudgetExceededError:` handlers keep
# working; custom callbacks reading StandardLoggingPayload pick these

View file

@ -382,15 +382,25 @@ class MCPClient:
if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError):
raise root_cause from in_flight_error
async def run_with_session(self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]) -> TSessionResult:
"""Open a session, run the provided coroutine, and clean up."""
async def run_with_session(
self,
operation: Callable[[ClientSession], Awaitable[TSessionResult]],
*,
quiet_on_error: bool = False,
) -> TSessionResult:
"""Open a session, run the provided coroutine, and clean up.
quiet_on_error demotes the failure line to debug for callers that own the exception
(call_tool / list_tools under raise_on_error), so an expected pass-through re-auth does
not emit a warning per call; every other caller keeps the operator-visible warning."""
http_client: Optional[httpx.AsyncClient] = None
try:
self._last_initialize_instructions = None
transport_ctx, http_client = self._create_transport_context()
return await self._execute_session_operation(transport_ctx, operation)
except Exception:
verbose_logger.warning("MCP client run_with_session failed for %s", self.server_url or "stdio")
_log = verbose_logger.debug if quiet_on_error else verbose_logger.warning
_log("MCP client run_with_session failed for %s", self.server_url or "stdio")
raise
finally:
if http_client is not None:
@ -491,7 +501,7 @@ class MCPClient:
return await session.list_tools()
try:
result = await self.run_with_session(_list_tools_operation)
result = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
tool_count = len(result.tools)
tool_names = [tool.name for tool in result.tools]
verbose_logger.info(f"MCP client listed {tool_count} tools from {self.server_url or 'stdio'}: {tool_names}")
@ -501,7 +511,13 @@ class MCPClient:
raise
except Exception as e:
error_type = type(e).__name__
verbose_logger.exception(
# Mirror call_tool: when the caller opted into raise_on_error it owns the exception and
# logs it at the fitting level (an expected pass-through re-auth 401 is info, not an
# error), so log at debug here to avoid an error-level line + traceback that would trip
# error-rate alerts on that expected signal. The swallow path still logs the full
# exception because nothing downstream will surface the failure.
_log = verbose_logger.debug if raise_on_error else verbose_logger.exception
_log(
f"MCP client list_tools failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
@ -510,7 +526,8 @@ class MCPClient:
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
_log_broken = verbose_logger.debug if raise_on_error else verbose_logger.error
_log_broken(
"MCP client detected broken connection/stream during list_tools - "
"the MCP server may have crashed, disconnected, or timed out"
)
@ -567,7 +584,7 @@ class MCPClient:
)
try:
tool_result = await self.run_with_session(_call_tool_operation)
tool_result = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error)
verbose_logger.info(f"MCP client tool call '{call_tool_request_params.name}' completed successfully")
return tool_result
except asyncio.CancelledError:
@ -580,7 +597,13 @@ class MCPClient:
verbose_logger.debug(f"MCP client tool call traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
# When the caller opted into raise_on_error it owns the exception and logs it at the
# level that fits (an expected pass-through re-auth 401 is info, not an operator-actionable
# error), so log at debug here to avoid an error-level line that would trip error-rate
# alerts on that expected signal. The swallow path (raise_on_error=False) still logs at
# error because nothing downstream will surface the failure.
_log = verbose_logger.debug if raise_on_error else verbose_logger.error
_log(
f"MCP client call_tool failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
@ -590,7 +613,7 @@ class MCPClient:
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
_log(
"MCP client detected broken connection/stream - "
"the MCP server may have crashed, disconnected, or timed out."
)

View file

@ -17,6 +17,7 @@ from litellm.llms.base_llm.google_genai.transformation import (
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes
from litellm.utils import ProviderConfigManager, client
if TYPE_CHECKING:
@ -39,6 +40,11 @@ base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
def _mark_async_entrypoint(logging_obj: LiteLLMLoggingObj | None, marker: str, is_async: bool) -> None:
if logging_obj is not None:
logging_obj.model_call_details.setdefault("litellm_params", {})[marker] = is_async
class GenerateContentSetupResult(BaseModel):
"""Internal Type - Result of setting up a generate content call"""
@ -315,6 +321,8 @@ def generate_content(
try:
_is_async = kwargs.pop("agenerate_content", False)
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content.value, _is_async)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")
@ -403,6 +411,8 @@ async def agenerate_content_stream(
try:
kwargs["agenerate_content_stream"] = True
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, True)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")
@ -497,6 +507,8 @@ def generate_content_stream(
# Remove any async-related flags since this is the sync function
_is_async = kwargs.pop("agenerate_content_stream", False)
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, _is_async)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")

View file

@ -1,3 +1,4 @@
import os
import secrets
from datetime import datetime
from typing import (
@ -17,6 +18,7 @@ from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.secret_managers.main import str_to_bool
from litellm.types.guardrails import (
DynamicGuardrailParams,
GuardrailEventHooks,
@ -59,6 +61,20 @@ from litellm.exceptions import (
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
def _strict_guardrail_modes_enabled() -> bool:
"""Whether guardrail-mode validation raises (default) or logs a warning.
Set `LITELLM_STRICT_GUARDRAIL_MODES=false` to keep the pre-LIT-4226 behavior
for guardrails whose supported_event_hooks list newly includes their
configured mode: log the mismatch and continue instead of raising at boot.
"""
raw = os.environ.get("LITELLM_STRICT_GUARDRAIL_MODES")
if raw is None:
return True
parsed = str_to_bool(raw)
return True if parsed is None else parsed
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
"""Extract session_id from request data (litellm_session_id or metadata)."""
session_id = request_data.get("litellm_session_id")
@ -132,7 +148,17 @@ class CustomGuardrail(CustomLogger):
if supported_event_hooks:
## validate event_hook is in supported_event_hooks
self._validate_event_hook(event_hook, supported_event_hooks)
try:
self._validate_event_hook(event_hook, supported_event_hooks)
except ValueError as validation_error:
if _strict_guardrail_modes_enabled():
raise
verbose_logger.warning(
"%s. LITELLM_STRICT_GUARDRAIL_MODES=false; continuing "
"with unsupported event_hook. Set the env var to true "
"(default) to enforce validation and fail at startup.",
validation_error,
)
super().__init__(**kwargs)
def render_violation_message(self, default: str, context: Optional[Dict[str, Any]] = None) -> str:
@ -303,6 +329,18 @@ class CustomGuardrail(CustomLogger):
"""
return None
@classmethod
def get_supported_event_hooks(cls) -> Optional[List[GuardrailEventHooks]]:
"""
Returns the event hooks this guardrail supports, for the UI to render.
Subclasses should override to return their supported hooks list. When a
subclass returns None, the endpoint omits it from the per-provider map
and the UI is expected to fall back to the global `supported_modes`
list client-side.
"""
return None
def _validate_event_hook(
self,
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]],
@ -477,6 +515,22 @@ class CustomGuardrail(CustomLogger):
return True
return False
def uses_apply_guardrail_interface(self) -> bool:
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
def _deployment_pre_call_target(self) -> "CustomLogger":
if not self.uses_apply_guardrail_interface():
return self
try:
from litellm.proxy.utils import unified_guardrail
except ImportError as e:
raise ImportError(
f"Guardrail {self.guardrail_name or type(self).__name__} implements apply_guardrail, which needs "
"the litellm proxy dependencies to run at the deployment level. "
"Install them with: pip install 'litellm[proxy]'"
) from e
return unified_guardrail
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
@ -495,7 +549,10 @@ class CustomGuardrail(CustomLogger):
# CHECK IF GUARDRAIL REJECTS THE REQUEST
if call_type == CallTypes.completion or call_type == CallTypes.acompletion:
result = await self.async_pre_call_hook(
target = self._deployment_pre_call_target()
if target is not self:
kwargs["guardrail_to_apply"] = self
result = await target.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id=kwargs.get("user_api_key_user_id"),
team_id=kwargs.get("user_api_key_team_id"),
@ -505,7 +562,7 @@ class CustomGuardrail(CustomLogger):
),
cache=dc,
data=kwargs,
call_type=call_type.value or "acompletion", # type: ignore
call_type="completion" if call_type == CallTypes.completion else "acompletion",
)
if result is not None and isinstance(result, dict):

View file

@ -19,7 +19,7 @@ import os
import time
import traceback
from datetime import datetime as datetimeObj
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Sequence, Union
import httpx
from httpx import Response
@ -50,6 +50,7 @@ from litellm.types.integrations.base_health_check import IntegrationHealthCheckS
from litellm.types.integrations.datadog import (
DD_ERRORS,
DD_MAX_BATCH_SIZE,
DD_MAX_PAYLOAD_SIZE_BYTES,
DataDogStatus,
DatadogInitParams,
DatadogPayload,
@ -384,8 +385,10 @@ class DataDogLogger(
async def _send_with_413_split(self, batch: List) -> List:
"""
Send a batch, halving any sub-batch that 413s (payload too large) and retrying the
halves, since Datadog enforces a 5MB uncompressed limit per request.
Send a batch, halving any sub-batch that exceeds Datadog's intake limits before
sending, and halving again on a 413 (payload too large) response, since Datadog
enforces a 5MB uncompressed limit per request. The proactive split avoids paying
a serialize + gzip + round trip for a payload the intake is guaranteed to reject.
A 413 surfaces as a raised MaskedHTTPStatusError (httpx raise_for_status), not a
returned response, so both paths are handled. A lone event that still 413s is
@ -398,6 +401,11 @@ class DataDogLogger(
chunk = pending.pop()
if not chunk:
continue
if len(chunk) > 1 and self._exceeds_intake_limits(chunk):
mid = len(chunk) // 2
pending.append(chunk[mid:])
pending.append(chunk[:mid])
continue
try:
response = await self.async_send_compressed_data(chunk)
except Exception as e:
@ -436,6 +444,21 @@ class DataDogLogger(
def _undelivered(chunk: List, pending: List[List]) -> List:
return chunk + [event for remaining in reversed(pending) for event in remaining]
@staticmethod
def _exceeds_intake_limits(chunk: Sequence[DatadogPayload]) -> bool:
"""
True when a chunk would breach Datadog's log intake limits: more than
DD_MAX_BATCH_SIZE events per payload, or a serialized size above
DD_MAX_PAYLOAD_SIZE_BYTES (held under Datadog's 5MB uncompressed cap so
the batch is split before the intake rejects it with a 413).
"""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
if len(chunk) > DD_MAX_BATCH_SIZE:
return True
payload_size_bytes = len(safe_dumps(chunk).encode("utf-8"))
return payload_size_bytes > DD_MAX_PAYLOAD_SIZE_BYTES
async def flush_queue(self):
if self.flush_lock is None:
return

View file

@ -267,29 +267,10 @@ class LangfuseOtelLogger(OpenTelemetry):
# If no keys, return default from env (likely logging to console or something else)
return OpenTelemetryConfig.from_env()
# Determine endpoint - default to US cloud
langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
if langfuse_host:
# If LANGFUSE_HOST is provided, construct OTEL endpoint from it
if not langfuse_host.startswith("http"):
langfuse_host = "https://" + langfuse_host
endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel"
verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}")
else:
# Default to US cloud endpoint
endpoint = LANGFUSE_CLOUD_US_ENDPOINT
verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
public_key=public_key, secret_key=secret_key
)
otlp_auth_headers = f"Authorization={auth_header}"
return OpenTelemetryConfig(
exporter="otlp_http",
endpoint=endpoint,
headers=otlp_auth_headers,
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=LangfuseOtelLogger._get_langfuse_otel_host(),
)
@staticmethod
@ -316,33 +297,36 @@ class LangfuseOtelLogger(OpenTelemetry):
"LANGFUSE_PUBLIC_KEY and LANGFUSE_SECRET_KEY must be set for Langfuse OpenTelemetry integration."
)
# Determine endpoint - default to US cloud
langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=LangfuseOtelLogger._get_langfuse_otel_host(),
)
@staticmethod
def _build_langfuse_otel_config(
public_key: str, secret_key: str, langfuse_host: Optional[str]
) -> "OpenTelemetryConfig":
"""
Builds an OTLP HTTP config pointing at the Langfuse OTEL endpoint for the
given host (US cloud when no host is provided), authorized with the given keys.
"""
if langfuse_host:
# If LANGFUSE_HOST is provided, construct OTEL endpoint from it
if not langfuse_host.startswith("http"):
langfuse_host = "https://" + langfuse_host
endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel"
normalized_host = langfuse_host if langfuse_host.startswith("http") else f"https://{langfuse_host}"
endpoint = f"{normalized_host.rstrip('/')}/api/public/otel"
verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}")
else:
# Default to US cloud endpoint
endpoint = LANGFUSE_CLOUD_US_ENDPOINT
verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
public_key=public_key, secret_key=secret_key
)
otlp_auth_headers = f"Authorization={auth_header}"
# Prevent modification of global env vars which causes leakage
# os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint
# os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers
return OpenTelemetryConfig(
exporter="otlp_http",
endpoint=endpoint,
headers=otlp_auth_headers,
headers=f"Authorization={auth_header}",
)
@staticmethod
@ -378,6 +362,29 @@ class LangfuseOtelLogger(OpenTelemetry):
return dynamic_headers
def construct_dynamic_otel_config(
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional["OpenTelemetryConfig"]:
"""
Build a full per-request OTLP config from team/key dynamic Langfuse credentials.
Key-scoped credentials must define the export target, not just the auth
headers: without this, a proxy with no global LANGFUSE_* env vars keeps its
init-time fallback exporter (console), so key-level langfuse_otel silently
never reaches Langfuse.
"""
public_key = standard_callback_dynamic_params.get("langfuse_public_key")
secret_key = standard_callback_dynamic_params.get("langfuse_secret_key")
if not public_key or not secret_key:
return None
langfuse_host = standard_callback_dynamic_params.get("langfuse_host") or self._get_langfuse_otel_host()
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=langfuse_host,
)
def create_litellm_proxy_request_started_span(
self,
start_time: datetime,

View file

@ -28,6 +28,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
parse_semconv_opt_in,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.services import ServiceLoggerPayload
from litellm.types.utils import (
@ -948,12 +949,22 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
Returns:
Tracer: The tracer to use for this request
"""
dynamic_config = self._get_dynamic_otel_config_from_kwargs(kwargs)
if dynamic_config is not None:
verbose_logger.debug(
"[OTEL DEBUG] Using DYNAMIC config tracer with endpoint: %s",
dynamic_config.endpoint,
)
return self._get_tracer_with_dynamic_config(dynamic_config)
dynamic_headers = self._get_dynamic_otel_headers_from_kwargs(kwargs)
if dynamic_headers is not None:
# Create spans using a temporary tracer with dynamic headers
tracer_to_use = self._get_tracer_with_dynamic_headers(dynamic_headers)
verbose_logger.debug("[OTEL DEBUG] Using DYNAMIC tracer with headers: %s", dynamic_headers)
verbose_logger.debug(
"[OTEL DEBUG] Using DYNAMIC tracer with headers: %s", redact_string(str(dynamic_headers))
)
else:
# For langfuse_otel without dynamic headers, create a provider with env var credentials
if hasattr(self, "callback_name") and self.callback_name == "langfuse_otel":
@ -989,6 +1000,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return dynamic_headers if dynamic_headers else None
def _get_dynamic_otel_config_from_kwargs(self, kwargs: dict) -> Optional[OpenTelemetryConfig]:
"""Extract a full dynamic exporter config from kwargs if available."""
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get(
"standard_callback_dynamic_params"
)
if not standard_callback_dynamic_params:
return None
return self.construct_dynamic_otel_config(standard_callback_dynamic_params=standard_callback_dynamic_params)
def _get_tracer_with_dynamic_config(self, dynamic_config: OpenTelemetryConfig):
"""Create (or reuse) a tracer whose exporter target comes from a per-request config."""
from opentelemetry.sdk.trace import TracerProvider
cache_key = f"dynamic_config:{dynamic_config.exporter}:{dynamic_config.endpoint}:{dynamic_config.headers}"
if cache_key in self._tracer_provider_cache:
return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME)
temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config))
temp_provider.add_span_processor(self._get_span_processor(config_override=dynamic_config))
self._tracer_provider_cache[cache_key] = temp_provider
return temp_provider.get_tracer(LITELLM_TRACER_NAME)
def _get_tracer_with_dynamic_headers(self, dynamic_headers: dict):
"""Create a temporary tracer with dynamic headers for this request only."""
from opentelemetry.sdk.trace import TracerProvider
@ -1020,6 +1057,19 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"""
return None
def construct_dynamic_otel_config(
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional[OpenTelemetryConfig]:
"""
Construct a full exporter config from standard callback dynamic params.
Override this when team/key dynamic params must control the export
target (exporter kind + endpoint), not just the request headers. When
this returns a config, it takes precedence over
construct_dynamic_otel_headers for the request.
"""
return None
#########################################################
# End of Team/Key Based Logging Control Flow
#########################################################
@ -2747,7 +2797,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug("OpenTelemetry: No parent context found, creating root span")
return None, None
def _get_span_processor(self, dynamic_headers: Optional[dict] = None):
def _get_span_processor(
self,
dynamic_headers: Optional[dict] = None,
config_override: Optional[OpenTelemetryConfig] = None,
):
from opentelemetry.sdk.trace.export import (
BatchSpanProcessor,
ConsoleSpanExporter,
@ -2755,40 +2809,45 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
SpanExporter,
)
otel_exporter = config_override.exporter if config_override else self.OTEL_EXPORTER
otel_endpoint = config_override.endpoint if config_override else self.OTEL_ENDPOINT
otel_headers = config_override.headers if config_override else self.OTEL_HEADERS
verbose_logger.debug(
"OpenTelemetry Logger, initializing span processor \nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
"OpenTelemetry Logger, initializing span processor \nexporter: %s\nendpoint: %s\nheaders: %s",
otel_exporter,
otel_endpoint,
redact_string(str(otel_headers)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(headers=dynamic_headers or self.OTEL_HEADERS)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(headers=dynamic_headers or otel_headers)
if dynamic_headers:
verbose_logger.debug(
"[OTEL DEBUG] Creating span processor with DYNAMIC headers: %s",
{k: v[:20] + "..." if len(str(v)) > 20 else v for k, v in _split_otel_headers.items()},
redact_string(str(_split_otel_headers)),
)
elif config_override:
verbose_logger.debug(
"[OTEL DEBUG] Creating span processor with DYNAMIC config, endpoint: %s",
otel_endpoint,
)
else:
verbose_logger.debug("[OTEL DEBUG] Creating span processor with GLOBAL headers")
if hasattr(self.OTEL_EXPORTER, "export"): # Check if it has the export method that SpanExporter requires
if hasattr(otel_exporter, "export"): # Check if it has the export method that SpanExporter requires
verbose_logger.debug(
"OpenTelemetry: intiializing SpanExporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return SimpleSpanProcessor(cast(SpanExporter, self.OTEL_EXPORTER))
return SimpleSpanProcessor(cast(SpanExporter, otel_exporter))
if self.OTEL_EXPORTER == "console":
if otel_exporter == "console":
verbose_logger.debug(
"OpenTelemetry: intiializing console exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return BatchSpanProcessor(ConsoleSpanExporter())
elif (
self.OTEL_EXPORTER == "otlp_http"
or self.OTEL_EXPORTER == "http/protobuf"
or self.OTEL_EXPORTER == "http/json"
):
elif otel_exporter == "otlp_http" or otel_exporter == "http/protobuf" or otel_exporter == "http/json":
try:
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterHTTP,
@ -2801,13 +2860,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug(
"OpenTelemetry: intiializing http exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "traces")
normalized_endpoint = self._normalize_otel_endpoint(otel_endpoint, "traces")
return BatchSpanProcessor(
OTLPSpanExporterHTTP(endpoint=normalized_endpoint, headers=_split_otel_headers),
)
elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc":
elif otel_exporter == "otlp_grpc" or otel_exporter == "grpc":
try:
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterGRPC,
@ -2820,16 +2879,16 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug(
"OpenTelemetry: intiializing grpc exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "traces")
normalized_endpoint = self._normalize_otel_endpoint(otel_endpoint, "traces")
return BatchSpanProcessor(
OTLPSpanExporterGRPC(endpoint=normalized_endpoint, headers=_split_otel_headers),
)
else:
verbose_logger.debug(
"OpenTelemetry: intiializing console exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return BatchSpanProcessor(ConsoleSpanExporter())
@ -2841,7 +2900,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"OpenTelemetry Logger, initializing log exporter \nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
redact_string(str(self.OTEL_HEADERS)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
@ -2928,7 +2987,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
redact_string(str(self.OTEL_HEADERS)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)

View file

@ -223,6 +223,15 @@ lives in [`plumbing/`](./plumbing):
readers/exporters receive them alongside the server metrics, and one is built
and registered as the global only when none is set (mirroring how V2 owns trace
export).
- [`events.py`](./plumbing/events.py) — GenAI client events. Gated on
`enable_events` (`LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS`), a failed LLM call
records the semconv `gen_ai.client.operation.exception` log event at severity
WARN, carrying `exception.type` / `exception.message` / `exception.stacktrace`
and correlated to the failed span through the trace and span ids. The
`LoggerProvider` is resolved like the meter provider, except that an explicit
`NoOpLoggerProvider` global is an operator opt-out that builds no recorder at
all. The deprecated `error.*` span attributes and the `exception` span event
are still stamped by the emitter for backwards compatibility.
### Adapter

View file

@ -18,6 +18,7 @@ from litellm.integrations.otel.model.payloads import (
ServiceSpanData,
SpanError,
)
from litellm.integrations.otel.plumbing.events import GenAIEventRecorder
from litellm.integrations.otel.plumbing.providers import to_otel_span_kind
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError
from litellm.integrations.otel.model.spans import (
@ -77,9 +78,11 @@ class SpanEmitter:
tracer: Tracer,
config: OpenTelemetryV2Config,
mappers: Sequence[AttributeMapper] | None = None,
event_recorder: GenAIEventRecorder | None = None,
) -> None:
self._tracer = tracer
self._config = config
self._event_recorder = event_recorder
# The mapper chain is the sole source of span attributes. When not
# passed in, resolve it from the config so there's one source of truth.
self._mappers: list[AttributeMapper] = (
@ -223,6 +226,14 @@ class SpanEmitter:
ExceptionEvent.NAME,
{ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},
)
if self._event_recorder is not None and role is SpanRole.LLM_CALL:
self._event_recorder.record_operation_exception(
span_context=span.get_span_context(),
error_type=error_type,
message=message,
stack_trace=error.stack_trace,
timestamp_ns=end_time_ns,
)
# On success leave the status UNSET (the semconv default) rather than
# forcing OK — that matches the FastAPI server span and avoids implying a
# span-level health signal litellm doesn't actually evaluate. Only a

View file

@ -6,6 +6,7 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Callable, Iterator, Mapping, Sequence, cast
from opentelemetry.context import Context, attach, get_current
from opentelemetry.sdk._logs import LoggerProvider
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import Span, Tracer, get_current_span, use_span
@ -40,14 +41,17 @@ from litellm.integrations.otel.model.payloads import (
is_mcp_list_tools,
is_mcp_tool_call,
)
from litellm.integrations.otel.plumbing.events import GenAIEventRecorder
from litellm.integrations.otel.plumbing.metrics import (
GenAIMetricRecorder,
create_genai_metrics,
)
from litellm.integrations.otel.plumbing.providers import (
build_tracer_provider,
get_event_logger,
get_meter,
get_tracer,
resolve_logger_provider,
resolve_meter_provider,
)
from litellm.integrations.otel.plumbing.routing import TenantTracerCache
@ -104,7 +108,7 @@ class OpenTelemetryV2(CustomLogger):
config: OpenTelemetryV2Config | None = None,
callback_name: str | None = None,
tracer_provider: TracerProvider | None = None,
logger_provider: Any | None = None, # reserved for OTel logs
logger_provider: LoggerProvider | None = None,
meter_provider: Any | None = None,
**kwargs: Any,
) -> None:
@ -117,7 +121,12 @@ class OpenTelemetryV2(CustomLogger):
self.tracer: Tracer = get_tracer(self._tracer_provider, LITELLM_TRACER_NAME)
self._metrics_recorder = self._init_metrics(meter_provider)
self._metric_filter_error_logged = False
self._emitter = SpanEmitter(self.tracer, self.config, mappers=resolve_mappers(self.config.mapper_names))
self._emitter = SpanEmitter(
self.tracer,
self.config,
mappers=resolve_mappers(self.config.mapper_names),
event_recorder=self._init_events(logger_provider),
)
self._tenant_tracers = TenantTracerCache(self.config, callback_name, LITELLM_TRACER_NAME)
self._open_llm_calls: "OrderedDict[str, _LLMCallSpan]" = OrderedDict()
self._init_otel_logger_on_litellm_proxy()
@ -136,6 +145,22 @@ class OpenTelemetryV2(CustomLogger):
meter = get_meter(provider, LITELLM_TRACER_NAME)
return GenAIMetricRecorder(create_genai_metrics(meter), self.callback_name)
def _init_events(self, logger_provider: LoggerProvider | None) -> "GenAIEventRecorder | None":
"""Create the GenAI event recorder when events are enabled, else ``None``.
``logger_provider`` is an explicit override (tests inject one); otherwise the
provider is resolved from the OTel global so an operator-configured logs
pipeline receives the events, building and registering one only when no
global provider is set. A ``None`` resolution means the operator opted out
of the logs signal, so no recorder is built.
"""
if not self.config.enable_events:
return None
provider = resolve_logger_provider(self.config, logger_provider)
if provider is None:
return None
return GenAIEventRecorder(get_event_logger(provider, LITELLM_TRACER_NAME))
# ====================================================================== #
# Proxy global registration
# ====================================================================== #

View file

@ -177,6 +177,19 @@ class ExceptionEvent:
NAME: Final = "exception"
TYPE: Final = "exception.type"
MESSAGE: Final = "exception.message"
STACKTRACE: Final = "exception.stacktrace"
class GenAIEvent:
"""GenAI semconv event names, from the GenAI registry's *events* section.
``gen_ai.client.operation.exception`` is defined as a log-based event
(severity WARN) carrying the ``exception.*`` trio, correlated to the failed
span via the trace/span ids — the semconv-compliant home for GenAI failure
details, unlike the deprecated ``error.message`` span attribute.
"""
OPERATION_EXCEPTION: Final = "gen_ai.client.operation.exception"
class Server:

View file

@ -0,0 +1,52 @@
"""GenAI client events: the ``gen_ai.client.operation.exception`` log event.
The GenAI semantic conventions define exception recording for client
operations as a log-based event (severity WARN) carrying the ``exception.*``
attribute trio, correlated to the failed span through the trace/span ids —
not as a span attribute or span event. This module owns building and
emitting that event; the exporter pipeline it rides is built in
:mod:`litellm.integrations.otel.plumbing.providers`.
"""
from dataclasses import dataclass
from opentelemetry._events import Event, EventLogger
from opentelemetry._logs.severity import SeverityNumber
from opentelemetry.trace import SpanContext
from litellm.integrations.otel.model.semconv import ExceptionEvent, GenAIEvent
@dataclass(frozen=True, slots=True)
class GenAIEventRecorder:
event_logger: EventLogger
def record_operation_exception(
self,
span_context: SpanContext,
error_type: str,
message: str,
stack_trace: str | None,
timestamp_ns: int | None,
) -> None:
# ``exception.type`` and ``exception.message`` are the semconv-required
# pair and always ride the event; only the recommended stacktrace is
# conditional on the payload carrying one.
stacktrace = ((ExceptionEvent.STACKTRACE, stack_trace),) if stack_trace else ()
self.event_logger.emit(
Event(
name=GenAIEvent.OPERATION_EXCEPTION,
timestamp=timestamp_ns,
trace_id=span_context.trace_id,
span_id=span_context.span_id,
trace_flags=span_context.trace_flags,
severity_number=SeverityNumber.WARN,
attributes=dict(
(
(ExceptionEvent.TYPE, error_type),
(ExceptionEvent.MESSAGE, message),
*stacktrace,
)
),
)
)

View file

@ -2,9 +2,20 @@
from typing import TYPE_CHECKING, Any, Callable, Iterable
from opentelemetry import baggage, metrics
from opentelemetry import _logs, baggage, metrics
from opentelemetry._events import EventLogger
from opentelemetry._logs import LoggerProvider, NoOpLoggerProvider
from opentelemetry.context import Context
from opentelemetry.metrics import MeterProvider, NoOpMeterProvider
from opentelemetry.sdk._events import EventLoggerProvider
from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider
from opentelemetry.sdk._logs.export import (
BatchLogRecordProcessor,
ConsoleLogExporter,
InMemoryLogExporter,
LogExporter,
SimpleLogRecordProcessor,
)
from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider
from opentelemetry.sdk.resources import Resource
from opentelemetry.sdk.trace import ReadableSpan, SpanProcessor, TracerProvider
@ -224,6 +235,112 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
def _otlp_logs_endpoint(endpoint: str | None) -> str | None:
"""Point an OTLP/HTTP base endpoint at the ``/v1/logs`` signal path.
The OTLP/HTTP exporter only appends ``/v1/logs`` when it reads
``OTEL_EXPORTER_OTLP_ENDPOINT`` itself; an explicitly passed endpoint is used
verbatim, so a base URL would POST to the root. Mirror ``_otlp_traces_endpoint``
for the logs signal (rewriting a sibling signal path when present).
"""
if not endpoint:
return endpoint
endpoint = endpoint.rstrip("/")
if endpoint.endswith("/v1/logs"):
return endpoint
for other_signal in ("/v1/traces", "/v1/metrics"):
if endpoint.endswith(other_signal):
return endpoint[: -len(other_signal)] + "/v1/logs"
return endpoint + "/v1/logs"
def build_log_exporter(config: OpenTelemetryV2Config) -> LogExporter:
"""Build a log exporter mirroring the exporter selection of the other signals.
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers;
``in_memory`` buffers for tests. Like GenAI metrics, events ride the
single-destination shorthand fields, not the multi-exporter ``exporters`` list.
"""
kind = (config.exporter or "console").lower()
if kind in ("in_memory", "inmemory", "memory"):
return InMemoryLogExporter()
if kind in ("otlp_http", "http", "http/protobuf", "http/json"):
from opentelemetry.exporter.otlp.proto.http._log_exporter import (
OTLPLogExporter as HTTPLogExporter,
)
return HTTPLogExporter(
endpoint=_otlp_logs_endpoint(config.endpoint),
headers=parse_headers(config.headers),
)
if kind in ("otlp_grpc", "grpc"):
try:
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import (
OTLPLogExporter as GRPCLogExporter,
)
except ImportError as exc:
raise ImportError(
"OpenTelemetry OTLP gRPC log exporter is not available. Install "
"`opentelemetry-exporter-otlp` and `grpcio` (or `litellm[grpc]`)."
) from exc
return GRPCLogExporter(endpoint=config.endpoint, headers=parse_headers(config.headers))
return ConsoleLogExporter()
def build_logger_provider(
config: OpenTelemetryV2Config,
log_exporter: LogExporter | None = None,
) -> SDKLoggerProvider:
"""Build the :class:`LoggerProvider` GenAI events export through.
``log_exporter`` is an explicit override (tests inject an
``InMemoryLogExporter``); otherwise the exporter is selected from the config's
exporter kind via :func:`build_log_exporter`. Console and in-memory exporters
get a Simple processor (synchronous export, which tests rely on), everything
else a Batch processor — the same split as span processing.
"""
exporter = log_exporter if log_exporter is not None else build_log_exporter(config)
provider = SDKLoggerProvider(resource=build_resource(config))
use_simple = isinstance(exporter, (ConsoleLogExporter, InMemoryLogExporter))
provider.add_log_record_processor(
SimpleLogRecordProcessor(exporter) if use_simple else BatchLogRecordProcessor(exporter)
)
return provider
def resolve_logger_provider(
config: OpenTelemetryV2Config,
logger_provider: SDKLoggerProvider | None = None,
) -> SDKLoggerProvider | None:
"""Resolve the :class:`LoggerProvider` GenAI events record through, or ``None``
when the operator has opted out of the logs signal.
Same resolution order as :func:`resolve_meter_provider`: an injected provider
wins (DI/tests); an operator-configured SDK global is reused so events ride
their pipeline; an explicit ``NoOpLoggerProvider`` global is an opt-out and
yields ``None``, so no event is ever built. Only the default placeholder
global makes V2 build a provider from the config and publish it as the global.
"""
if logger_provider is not None:
return logger_provider
existing: LoggerProvider = _logs.get_logger_provider()
if isinstance(existing, SDKLoggerProvider):
return existing
if isinstance(existing, NoOpLoggerProvider):
return None
provider = build_logger_provider(config)
_logs.set_logger_provider(provider)
return provider
def get_event_logger(provider: SDKLoggerProvider, name: str = "litellm") -> EventLogger:
return EventLoggerProvider(logger_provider=provider).get_event_logger(name, litellm_version)
def build_meter_provider(
config: OpenTelemetryV2Config,
metric_reader: "MetricReader | None" = None,

View file

@ -239,6 +239,18 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_output_audio_tokens_metric"),
)
self.litellm_video_duration_seconds_metric = self._counter_factory(
"litellm_video_duration_seconds_metric",
"Seconds of video generated, from usage.duration_seconds on video generation calls",
labelnames=self.get_labels_for_metric("litellm_video_duration_seconds_metric"),
)
self.litellm_images_generated_metric = self._counter_factory(
"litellm_images_generated_metric",
"Number of images generated, from the image generation response",
labelnames=self.get_labels_for_metric("litellm_images_generated_metric"),
)
# Remaining Budget for Team
self.litellm_remaining_team_budget_metric = self._gauge_factory(
"litellm_remaining_team_budget_metric",
@ -1336,6 +1348,12 @@ class PrometheusLogger(CustomLogger):
label_context=label_context,
)
self._increment_media_generation_metrics(
standard_logging_payload=standard_logging_payload,
enum_values=enum_values,
label_context=label_context,
)
# MCP tool call metrics
self._increment_mcp_tool_call_metrics(
standard_logging_payload=standard_logging_payload,
@ -1459,8 +1477,65 @@ class PrometheusLogger(CustomLogger):
),
]
for counter, metric_name, value in detail_metrics:
if not isinstance(value, (int, float)) or value <= 0:
PrometheusLogger._inc_sparse_usage_counters(
self,
detail_metrics,
enum_values=enum_values,
label_context=label_context,
)
def _increment_media_generation_metrics(
self,
standard_logging_payload: StandardLoggingPayload,
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext | None = None,
) -> None:
"""
Increment video-seconds and images-generated counters from
``standard_logging_payload["metadata"]["usage_object"]``. Video
providers report ``duration_seconds`` there; image generation calls
report ``output_image_count``. Both are sparse: only emitted when the
value is present and > 0, so token-only call types are unaffected.
"""
metadata = standard_logging_payload.get("metadata") or {}
usage_object = metadata.get("usage_object") if isinstance(metadata, dict) else None
if not isinstance(usage_object, dict):
return
media_metrics: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
(
self.litellm_video_duration_seconds_metric,
"litellm_video_duration_seconds_metric",
usage_object.get("duration_seconds"),
),
(
self.litellm_images_generated_metric,
"litellm_images_generated_metric",
usage_object.get("output_image_count"),
),
]
PrometheusLogger._inc_sparse_usage_counters(
self,
media_metrics,
enum_values=enum_values,
label_context=label_context,
)
def _inc_sparse_usage_counters(
self,
counters_with_values: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]],
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext | None = None,
) -> None:
"""
Increment each ``(counter, metric_name, value)`` entry whose value is
a positive number. Non-numeric values (including booleans from
malformed provider usage dicts) and values <= 0 are skipped, keeping
scrape output sparse.
"""
for counter, metric_name, value in counters_with_values:
if isinstance(value, bool) or not isinstance(value, (int, float)) or value <= 0:
continue
PrometheusLogger._inc_labeled_counter(
self,
@ -1618,6 +1693,14 @@ class PrometheusLogger(CustomLogger):
user_id: Optional[str] = None,
user_api_key_org_id: Optional[str] = None,
):
if (
isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric)
):
return
_metadata = litellm_params.get("metadata") or {}
_team_spend = _metadata.get("user_api_key_team_spend", None)
_team_max_budget = _metadata.get("user_api_key_team_max_budget", None)
@ -1708,6 +1791,35 @@ class PrometheusLogger(CustomLogger):
amount=float(response_cost),
)
@staticmethod
def _get_remaining_from_v3_rate_limit_headers(
standard_logging_payload: StandardLoggingPayload | None,
rate_limit_type: Literal["requests", "tokens"],
) -> int | None:
"""
Read the per-(key, model) remaining value emitted by the v3 rate
limiter (``parallel_request_limiter_v3.py``), which writes
``x-ratelimit-model_per_key-remaining-{requests,tokens}`` into
``standard_logging_object.hidden_params.additional_headers`` instead
of the ``litellm-key-remaining-*`` metadata keys the legacy limiter
sets. The header carries no model group; it always refers to this
request's model group, which is what the gauges are labeled with.
Values are written in-process as plain ints (never HTTP-serialized
strings), so anything else is rejected rather than coerced.
"""
if standard_logging_payload is None:
return None
hidden_params = standard_logging_payload.get("hidden_params")
if hidden_params is None:
return None
additional_headers = hidden_params.get("additional_headers")
if additional_headers is None:
return None
value = dict(additional_headers).get(f"x-ratelimit-model_per_key-remaining-{rate_limit_type}")
if isinstance(value, bool) or not isinstance(value, int):
return None
return value
def _set_virtual_key_rate_limit_metrics(
self,
user_api_key: Optional[str],
@ -1725,11 +1837,20 @@ class PrometheusLogger(CustomLogger):
model_group = get_model_group_from_litellm_kwargs(kwargs)
remaining_requests_variable_name = f"litellm-key-remaining-requests-{model_group}"
remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}"
standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
remaining_requests = metadata.get(remaining_requests_variable_name)
if remaining_requests is None:
remaining_requests = self._get_remaining_from_v3_rate_limit_headers(
standard_logging_payload=standard_logging_payload, rate_limit_type="requests"
)
if remaining_requests is None:
remaining_requests = sys.maxsize
remaining_tokens = metadata.get(remaining_tokens_variable_name)
if remaining_tokens is None:
remaining_tokens = self._get_remaining_from_v3_rate_limit_headers(
standard_logging_payload=standard_logging_payload, rate_limit_type="tokens"
)
if remaining_tokens is None:
remaining_tokens = sys.maxsize
@ -3332,6 +3453,9 @@ class PrometheusLogger(CustomLogger):
- looks up team info from db if not available in metadata
- Set team budget metrics
"""
if isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric):
return
if user_api_team:
team_object = await self._assemble_team_object(
team_id=user_api_team,
@ -3453,6 +3577,9 @@ class PrometheusLogger(CustomLogger):
- Fetches org info via cache (get_org_object)
- Sets org budget metrics
"""
if isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric):
return
if not org_id:
return
@ -3582,6 +3709,9 @@ class PrometheusLogger(CustomLogger):
key_max_budget: Optional[float],
key_spend: Optional[float],
):
if isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric):
return
if user_api_key:
user_api_key_dict = await self._assemble_key_object(
user_api_key=user_api_key,
@ -3642,6 +3772,9 @@ class PrometheusLogger(CustomLogger):
- looks up user info from db if not available in metadata
- Set user budget metrics
"""
if isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric):
return
if user_id:
user_object = await self._assemble_user_object(
user_id=user_id,

View file

@ -7,7 +7,7 @@ import time
import urllib.parse
import uuid
from collections import Counter
from typing import TYPE_CHECKING, Any, Literal, Optional
from typing import TYPE_CHECKING, Any, List, Literal, Optional
import httpx
from litellm._logging import verbose_logger
@ -52,6 +52,10 @@ class _MalformedToolBlockingResponseError(Exception):
class RubrikLogger(CustomGuardrail, CustomBatchLogger):
@classmethod
def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]:
return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]
def __init__(
self,
api_key: str | None = None,
@ -69,6 +73,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
kwargs["event_hook"] = kwargs.get("event_hook") or GuardrailEventHooks.post_call
if kwargs.get("default_on") is None:
kwargs["default_on"] = True
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(
flush_lock=self.flush_lock,
**kwargs,

View file

@ -161,8 +161,13 @@ def get_s3_object_key(
start_time: datetime,
s3_file_name: str,
) -> str:
sanitized_s3_file_name = s3_file_name.replace("/", "_")
s3_object_key = (
(s3_path.rstrip("/") + "/" if s3_path else "") + prefix + start_time.strftime("%Y-%m-%d") + "/" + s3_file_name
(s3_path.rstrip("/") + "/" if s3_path else "")
+ prefix
+ start_time.strftime("%Y-%m-%d")
+ "/"
+ sanitized_s3_file_name
) # we need the s3 key to include the time, so we log cache hits too
s3_object_key += ".json"
return s3_object_key

View file

@ -19,9 +19,11 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.websearch_interception.tools import (
get_litellm_web_search_tool,
get_litellm_web_search_tool_openai,
get_litellm_web_search_tool_responses,
is_anthropic_native_web_search_tool,
is_web_search_tool,
is_web_search_tool_chat_completion,
is_web_search_tool_responses,
)
from litellm.integrations.websearch_interception.transformation import (
WebSearchTransformation,
@ -32,11 +34,12 @@ from litellm.types.integrations.websearch_interception import (
)
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
RESPONSES_AGENTIC_SURFACE,
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import LlmProviders
from litellm.types.utils import CallTypes, LlmProviders
from litellm.utils import ProviderConfigManager
# Key used to flag, on per-request kwargs, that the originating client sent
@ -251,6 +254,9 @@ class WebSearchInterceptionLogger(CustomLogger):
if not tools:
return None
if call_type in (CallTypes.responses, CallTypes.aresponses):
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
# Check if any tool is a web search tool (native or already LiteLLM standard)
has_websearch = any(is_web_search_tool(t) for t in tools)
@ -291,6 +297,26 @@ class WebSearchInterceptionLogger(CustomLogger):
return kwargs
def _convert_responses_tools(self, kwargs: dict[str, Any], tools: list[dict[str, Any]]) -> dict | None:
"""Convert Responses API web search tools to the LiteLLM standard function tool."""
if not any(is_web_search_tool_responses(tool) for tool in tools):
return None
verbose_logger.debug("WebSearchInterception: Converting Responses web_search tools to LiteLLM standard")
converted_tools = [
get_litellm_web_search_tool_responses() if is_web_search_tool_responses(tool) else tool for tool in tools
]
converted_kwargs = {**kwargs, "tools": converted_tools}
if kwargs.get("stream"):
verbose_logger.debug("WebSearchInterception: deployment hook converting stream=True to stream=False")
converted_kwargs["stream"] = False
converted_kwargs["_websearch_interception_converted_stream"] = True
return converted_kwargs
@classmethod
def from_config_yaml(cls, config: WebSearchInterceptionConfig) -> "WebSearchInterceptionLogger":
"""
@ -461,6 +487,17 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs=kwargs,
)
if kwargs.get("_agentic_loop_api_surface") == RESPONSES_AGENTIC_SURFACE:
return await self.async_should_run_responses_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
verbose_logger.debug(f"WebSearchInterception: Hook called! provider={custom_llm_provider}, stream={stream}")
verbose_logger.debug(f"WebSearchInterception: Response type: {type(response)}")
@ -597,6 +634,54 @@ class WebSearchInterceptionLogger(CustomLogger):
}
return True, tools_dict
async def async_should_run_responses_agentic_loop(
self,
response: Any,
model: str,
messages: list[dict],
tools: list[dict] | None,
stream: bool,
custom_llm_provider: str,
kwargs: dict,
) -> tuple[bool, dict]:
"""Check if WebSearch interception is needed for the Responses API."""
verbose_logger.debug(
f"WebSearchInterception: Responses hook called! provider={custom_llm_provider}, stream={stream}"
)
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
verbose_logger.debug(
f"WebSearchInterception: Skipping provider {custom_llm_provider} (not in enabled list: {self.enabled_providers})"
)
return False, {}
has_websearch_tool = any(is_web_search_tool_responses(t) for t in (tools or []))
if not has_websearch_tool:
verbose_logger.debug("WebSearchInterception: No litellm_web_search tool in responses request")
return False, {}
should_intercept, tool_calls = WebSearchTransformation.transform_request(
response=response,
stream=stream,
response_format="responses",
)
if not should_intercept:
verbose_logger.debug("WebSearchInterception: No WebSearch function_call detected in responses output")
return False, {}
verbose_logger.debug(
f"WebSearchInterception: Detected {len(tool_calls)} WebSearch function_call(s), executing agentic loop"
)
tools_dict = {
"tool_calls": tool_calls,
"tool_type": "websearch",
"provider": custom_llm_provider,
"response_format": "responses",
}
return True, tools_dict
async def async_run_agentic_loop(
self,
tools: Dict,
@ -655,6 +740,18 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs=kwargs,
)
if kwargs.get("_agentic_loop_api_surface") == RESPONSES_AGENTIC_SURFACE:
return await self.async_build_responses_agentic_loop_plan(
tools=tools,
model=model,
messages=messages,
response=response,
optional_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs,
)
tool_calls = tools["tool_calls"]
thinking_blocks = tools.get("thinking_blocks", [])
request_patch, structured_results = await self._build_anthropic_request_patch(
@ -809,6 +906,133 @@ class WebSearchInterceptionLogger(CustomLogger):
metadata={"tool_type": "websearch", "response_format": response_format},
)
async def async_build_responses_agentic_loop_plan(
self,
tools: dict,
model: str,
messages: list[dict],
response: Any,
optional_params: dict,
logging_obj: Any,
stream: bool,
kwargs: dict,
) -> AgenticLoopPlan:
tool_calls = tools["tool_calls"]
request_patch = await self._build_responses_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
optional_params=optional_params,
kwargs=kwargs,
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "websearch", "response_format": "responses"},
)
async def _build_responses_request_patch(
self,
model: str,
messages: Union[str, list[dict]],
tool_calls: list[dict],
optional_params: dict,
kwargs: dict,
) -> AgenticLoopRequestPatch:
"""Execute litellm.asearch() and build a Responses API rerun patch."""
search_tasks = [
(
self._execute_search(tool_call["input"]["query"], kwargs=kwargs)
if isinstance(tool_call.get("input"), dict) and tool_call["input"].get("query")
else self._create_empty_search_result()
)
for tool_call in tool_calls
]
verbose_logger.debug(f"WebSearchInterception: Executing {len(search_tasks)} responses search(es) in parallel")
search_results = await asyncio.gather(*search_tasks, return_exceptions=True)
search_texts = [self._extract_search_text(result) for result in search_results]
followup_items = [
item
for tool_call, search_text in zip(tool_calls, search_texts)
for item in (
{
"type": "function_call",
"call_id": tool_call.get("call_id"),
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
"arguments": tool_call.get("arguments", ""),
},
{
"type": "function_call_output",
"call_id": tool_call.get("call_id"),
"output": search_text,
},
)
]
input_list = self._normalize_responses_input(messages) + followup_items
tools_param = optional_params.get("tools")
optional_params_clean = {
k: v
for k, v in optional_params.items()
if k not in {"tools", "tool_choice", "stream", "model_alias_map", "stream_response", "custom_prompt_dict"}
}
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and k
not in {
"_agentic_loop_api_surface",
"litellm_logging_obj",
"acompletion",
"custom_llm_provider",
"model_alias_map",
}
}
full_model_name = model
if "/" not in model and isinstance(kwargs.get("custom_llm_provider"), str):
full_model_name = f"{kwargs['custom_llm_provider']}/{model}"
verbose_logger.debug(
"WebSearchInterception: Built responses request patch model=%s input_items=%d searches=%d",
full_model_name,
len(input_list),
len(search_texts),
)
return AgenticLoopRequestPatch(
model=full_model_name,
messages=input_list,
tools=tools_param if isinstance(tools_param, list) else None,
optional_params=optional_params_clean,
kwargs=kwargs_for_followup,
)
@staticmethod
def _normalize_responses_input(messages: Union[str, list[dict]]) -> list[dict]:
if isinstance(messages, str):
return [{"role": "user", "content": messages}]
if isinstance(messages, list):
return list(messages)
return []
@staticmethod
def _extract_search_text(result: Any) -> str:
if isinstance(result, Exception):
verbose_logger.error(f"WebSearchInterception: Responses search failed with error: {str(result)}")
return f"Search failed: {str(result)}"
if isinstance(result, tuple) and len(result) == 2:
text_value, _ = result
return text_value if isinstance(text_value, str) else str(text_value)
verbose_logger.debug(f"WebSearchInterception: Unexpected search result type {type(result)}")
return str(result)
@staticmethod
def _resolve_max_tokens(
optional_params: Dict,

View file

@ -82,6 +82,75 @@ def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
}
def get_litellm_web_search_tool_responses() -> dict[str, Any]:
"""
Get the standard LiteLLM web search tool definition in Responses API format.
Used by async_pre_call_deployment_hook on the Responses API path, where a
function tool is a flat object (``type: "function"`` with a top-level
``name`` and ``parameters``) rather than the nested ``function`` wrapper
used by Chat Completions.
Returns:
Dict containing the Responses-style function tool definition.
"""
return {
"type": "function",
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
"description": (
"Search the web for information. Use this when you need current "
"information or answers to questions that require up-to-date data."
),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "The search query to execute",
}
},
"required": ["query"],
},
}
def is_web_search_tool_responses(tool: dict[str, Any]) -> bool:
"""
Check if a tool is a web search tool for the Responses API.
Detects:
- OpenAI native Responses web search tools, whose ``type`` is one of
``web_search``, ``web_search_2025_08_26``, ``web_search_preview``,
``web_search_preview_2025_03_11`` (matched by the ``web_search`` prefix)
- The LiteLLM standard function tool in Responses shape:
``{"type": "function", "name": "litellm_web_search"}``
Args:
tool: Tool dictionary to check
Returns:
True if tool is a Responses-API web search tool
Example:
>>> is_web_search_tool_responses({"type": "web_search"})
True
>>> is_web_search_tool_responses({"type": "web_search_preview"})
True
>>> is_web_search_tool_responses({"type": "function", "name": "litellm_web_search"})
True
>>> is_web_search_tool_responses({"type": "function", "name": "get_weather"})
False
"""
tool_type = tool.get("type", "")
if not isinstance(tool_type, str):
return False
if tool_type == "function":
return tool.get("name") == LITELLM_WEB_SEARCH_TOOL_NAME
return tool_type == "web_search" or tool_type.startswith("web_search_")
def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool:
"""
Check if a tool is a web search tool for Chat Completions API (strict check).

View file

@ -59,9 +59,73 @@ class WebSearchTransformation:
# Parse non-streaming response based on format
if response_format == "openai":
return WebSearchTransformation._detect_from_openai_response(response)
elif response_format == "responses":
return WebSearchTransformation._detect_from_responses_response(response)
else:
return WebSearchTransformation._detect_from_non_streaming_response(response)
@staticmethod
def _detect_from_responses_response(
response: Any,
) -> tuple[bool, list[dict]]:
"""Parse a Responses API response for ``litellm_web_search`` function calls.
After pre-request conversion the native web search tool is replaced by a
``litellm_web_search`` function tool, so the model emits ``function_call``
items in ``response.output`` instead of a native ``web_search_call``.
"""
if isinstance(response, dict):
output = response.get("output", [])
else:
output = getattr(response, "output", None) or []
if not isinstance(output, list):
return False, []
tool_calls: list[dict] = []
for item in output:
if isinstance(item, dict):
item_type = item.get("type")
item_name = item.get("name")
call_id = item.get("call_id")
arguments = item.get("arguments", "")
else:
item_type = getattr(item, "type", None)
item_name = getattr(item, "name", None)
call_id = getattr(item, "call_id", None)
arguments = getattr(item, "arguments", "")
if item_type != "function_call" or item_name != LITELLM_WEB_SEARCH_TOOL_NAME:
continue
if isinstance(arguments, str):
try:
parsed_input = json.loads(arguments) if arguments else {}
except json.JSONDecodeError:
verbose_logger.warning(
f"WebSearchInterception: Failed to parse function_call arguments: {arguments}"
)
parsed_input = {}
elif isinstance(arguments, dict):
parsed_input = arguments
else:
parsed_input = {}
arguments_str = arguments if isinstance(arguments, str) else json.dumps(parsed_input)
tool_calls.append(
{
"id": call_id,
"call_id": call_id,
"type": "function_call",
"name": item_name,
"arguments": arguments_str,
"input": parsed_input,
}
)
verbose_logger.debug(f"WebSearchInterception: Found {item_name} function_call with call_id={call_id}")
return len(tool_calls) > 0, tool_calls
@staticmethod
def _detect_from_non_streaming_response(
response: Any,

View file

@ -7,6 +7,7 @@ This module has no dependencies on proxy code and can be safely imported at the
import json
import os
import time
from pathlib import Path
from typing import Optional
@ -68,3 +69,17 @@ def get_litellm_gateway_api_key(
if stored_url != expected_base_url.rstrip("/"):
return None
return token_data["key"]
def is_cli_token_fresh(token_data: dict, buffer_hours: float = 0.1) -> bool:
"""Check whether a cached CLI token (as stored in token.json) is still
within its expiration window. Used by `lite auth print-token` to fail
fast, without a network round trip, once the cached token is past
`LITELLM_CLI_JWT_EXPIRATION_HOURS`."""
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
timestamp = token_data.get("timestamp")
if not isinstance(timestamp, (int, float)):
return False
age_hours = (time.time() - timestamp) / 3600
return age_hours < (CLI_JWT_EXPIRATION_HOURS - buffer_hours)

View file

@ -10,7 +10,7 @@ from typing import TYPE_CHECKING, Any, Optional, Union
from litellm.secret_managers.main import get_secret_bool
if TYPE_CHECKING:
from ddtrace.tracer import Tracer as DD_TRACER
from ddtrace.trace import Tracer as DD_TRACER
else:
DD_TRACER = Any

View file

@ -19,3 +19,19 @@ def get_env_int(env_var: str, default: int) -> int:
return int(raw)
except (ValueError, TypeError):
return default
def get_env_int_or_none(env_var: str) -> int | None:
"""Parse an environment variable as an integer, returning None when it is unset or unusable.
Use this instead of `get_env_int` when callers must distinguish "explicitly configured"
from "left at the default", for example when an override should take precedence over a
value resolved from somewhere else.
"""
raw = os.getenv(env_var)
if raw is None:
return None
try:
return int(raw.strip())
except (ValueError, TypeError):
return None

View file

@ -95,6 +95,9 @@ class ExceptionCheckers:
if "current length is" in _error_str_lowercase and "while limit is" in _error_str_lowercase:
return True
if "maximum input length is" in _error_str_lowercase and "tokens" in _error_str_lowercase:
return True
return False
@staticmethod

View file

@ -3,52 +3,69 @@ Declarative fallback generalizations for unknown / newly-released models.
The ``fallback_generalizations`` block in ``model_prices_and_context_window.json``
holds an ordered list of rules. Each rule pairs a single case-insensitive regex
with the metadata to apply when a model name has no exact entry in the cost map.
The metadata is a partial cost-map entry: ``litellm_provider`` drives provider
routing, and the remaining fields (``mode``, ``supports_*``, context window,
pricing, ...) drive ``get_model_info`` / ``supports_*``.
with a ``model_info`` dict, and the structure of ``model_info`` decides which of
two kinds the rule is.
Precedence: rules are evaluated in file order and the first match wins. They are
consulted only after exact and case-insensitive lookups miss, so an exact entry
always takes precedence over a rule.
A ROUTING rule carries exactly one ``model_info`` key, ``litellm_provider``. It is
consumed only by ``get_llm_provider`` bare-id inference: the first routing rule
whose regex matches decides the provider. Routing rules never contribute to model
info.
A CAPABILITY rule carries any ``model_info`` keys except ``litellm_provider``
(``mode``, ``supports_*``, context window, pricing, ...). It is consumed by
``get_model_info`` fallback resolution: the ``model_info`` of ALL capability rules
whose regex matches is unioned in file order, with later rules overriding earlier
ones on key conflicts, and the caller backfills ``litellm_provider`` with the
provider it requested. If no capability rule matches, model-info resolution misses
as if no rules existed.
LEGACY-SCHEMA SHIM (temporary, until the new-schema JSON reaches main): released
proxies fetch this JSON remotely from main, whose block still ships the old schema
where a rule mixes ``litellm_provider`` with capability keys and may inherit a
parent's ``model_info`` via ``extends``. Such a legacy rule is tolerated rather
than skipped: ``extends`` is resolved once at install time (single level, against
raw parents), and the resolved rule acts as BOTH kinds, a routing rule (its
``litellm_provider`` participates in first-hit inference) and a capability rule
(its full ``model_info``, provider included, participates in the union). New-schema
rules never mix the two and never use ``extends``. A rule whose
``litellm_provider`` is not a string is invalid and is warned about and skipped
(a warning rather than a crash, for the same remote-fetch reason).
Rules are only consulted after exact and case-insensitive lookups miss, so an
exact cost-map entry always takes precedence over any rule.
Patterns are matched case-insensitively with ``re.search`` and are not implicitly
anchored: a rule must include ``^`` and ``$`` (as the shipped rules do) to bind to
the whole model name, otherwise it matches as a substring. Keeping anchoring in the
regex makes the rule the single, self-contained source of truth for what it matches.
A rule may set ``extends`` to the ``name`` of another rule to inherit that rule's
``model_info``; the rule's own ``model_info`` overrides the inherited keys, so a
narrow rule (for example a version-gated capability flag) carries only its delta
instead of duplicating the parent's pricing block. Inheritance is resolved once,
at install time, against each rule's raw (unresolved) ``model_info``; it is a
single level (a parent that itself extends is not chained).
anchored: a rule must include ``^`` and ``$`` to bind to the whole model name,
otherwise it matches as a substring. Keeping anchoring in the regex makes the rule
the single, self-contained source of truth for what it matches.
Any other keys on a rule (for example a free-text ``description`` documenting what
the regex matches) are ignored by the engine and exist only for the reader.
The compiled-regex list is built once and cached. ``match_fallback_generalization``
is O(number of rules); callers must only invoke it on a cache miss.
Rules are compiled and classified once, at install time. The match functions are
O(number of rules); callers must only invoke them on a cache miss.
"""
import re
from typing import Optional
from dataclasses import dataclass
from typing import Optional, Union
from litellm._logging import verbose_logger
NAME_FIELD = "name"
PATTERN_FIELD = "pattern"
MODEL_INFO_FIELD = "model_info"
EXTENDS_FIELD = "extends"
PROVIDER_KEY = "litellm_provider"
LEGACY_EXTENDS_FIELD = "extends"
def _resolve_extends(rules: list) -> list:
"""Expand ``extends`` inheritance so each rule's ``model_info`` is self-contained.
def _resolve_legacy_extends(rules: list) -> list:
"""Expand legacy ``extends`` inheritance so each rule's ``model_info`` is self-contained.
A rule with ``extends: <name>`` is rewritten with ``model_info`` set to the parent's
``model_info`` overlaid by its own. Resolution is single-level and uses each rule's
raw ``model_info`` as the parent source. Non-dict rules and dangling parents are
passed through unchanged.
Compatibility shim for the old remote schema: single level, resolved against each
parent's raw ``model_info``, with the child's own keys winning on conflict. Non-dict
rules and dangling parents pass through unchanged; new-schema rules carry no
``extends`` and are untouched.
"""
base_by_name = {
rule[NAME_FIELD]: rule[MODEL_INFO_FIELD]
@ -58,84 +75,138 @@ def _resolve_extends(rules: list) -> list:
and isinstance(rule.get(MODEL_INFO_FIELD), dict)
}
def resolved(rule: dict) -> dict:
parent_name = rule.get(EXTENDS_FIELD)
def resolved(rule: object) -> object:
if not isinstance(rule, dict):
return rule
parent_name = rule.get(LEGACY_EXTENDS_FIELD)
own_info = rule.get(MODEL_INFO_FIELD)
parent_info = base_by_name.get(parent_name) if isinstance(parent_name, str) else None
if parent_info is None or not isinstance(own_info, dict):
return rule
return {**rule, MODEL_INFO_FIELD: {**parent_info, **own_info}}
return [resolved(rule) if isinstance(rule, dict) else rule for rule in rules]
return [resolved(rule) for rule in rules]
@dataclass(frozen=True, slots=True)
class _RoutingRule:
pattern: re.Pattern
provider: str
@dataclass(frozen=True, slots=True)
class _CapabilityRule:
pattern: re.Pattern
model_info: dict
_CompiledRule = Union[_RoutingRule, _CapabilityRule]
def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:
if not isinstance(rule, dict):
return ()
pattern = rule.get(PATTERN_FIELD)
model_info = rule.get(MODEL_INFO_FIELD)
if not isinstance(pattern, str) or not isinstance(model_info, dict):
verbose_logger.warning(
"LiteLLM: skipping malformed fallback generalization rule %s (needs string '%s' and dict '%s').",
rule.get(NAME_FIELD, pattern),
PATTERN_FIELD,
MODEL_INFO_FIELD,
)
return ()
try:
compiled = re.compile(pattern, re.IGNORECASE)
except re.error as e:
verbose_logger.warning(
"LiteLLM: skipping fallback generalization rule with invalid regex %r: %s",
pattern,
e,
)
return ()
if PROVIDER_KEY not in model_info:
return (_CapabilityRule(pattern=compiled, model_info=model_info),)
provider = model_info[PROVIDER_KEY]
if not isinstance(provider, str):
verbose_logger.warning(
"LiteLLM: skipping invalid fallback generalization rule %s: '%s' in '%s' must be a string.",
rule.get(NAME_FIELD, pattern),
PROVIDER_KEY,
MODEL_INFO_FIELD,
)
return ()
if len(model_info) == 1:
return (_RoutingRule(pattern=compiled, provider=provider),)
return (
_RoutingRule(pattern=compiled, provider=provider),
_CapabilityRule(pattern=compiled, model_info=model_info),
)
class _FallbackGeneralizations:
"""Holds the active rule list and its lazily-compiled regex cache."""
"""Holds the raw rule list and its install-time-compiled routing and capability rules."""
def __init__(self) -> None:
self.rules: list[dict] = []
self._compiled: Optional[list[tuple[re.Pattern, dict]]] = None
self.rules: list = []
self.routing_rules: tuple = ()
self.capability_rules: tuple = ()
def set_rules(self, rules: Optional[list[dict]]) -> None:
self.rules = rules if isinstance(rules, list) else []
self._compiled = None
def set_rules(self, rules: Optional[list]) -> None:
installed = rules if isinstance(rules, list) else []
compiled = tuple(kind for rule in _resolve_legacy_extends(installed) for kind in _compile_rule(rule))
self.rules = installed
self.routing_rules = tuple(rule for rule in compiled if isinstance(rule, _RoutingRule))
self.capability_rules = tuple(rule for rule in compiled if isinstance(rule, _CapabilityRule))
def _compile(self) -> list[tuple[re.Pattern, dict]]:
compiled: list[tuple[re.Pattern, dict]] = []
for rule in self.rules:
if not isinstance(rule, dict):
continue
pattern = rule.get(PATTERN_FIELD)
model_info = rule.get(MODEL_INFO_FIELD)
if not isinstance(pattern, str) or not isinstance(model_info, dict):
verbose_logger.warning(
"LiteLLM: skipping malformed fallback generalization rule %s (needs string '%s' and dict '%s').",
rule.get("name", pattern),
PATTERN_FIELD,
MODEL_INFO_FIELD,
)
continue
try:
compiled.append((re.compile(pattern, re.IGNORECASE), model_info))
except re.error as e:
verbose_logger.warning(
"LiteLLM: skipping fallback generalization rule with invalid regex %r: %s",
pattern,
e,
)
return compiled
def match(self, model: str) -> Optional[dict]:
def match_routing(self, model: str) -> Optional[str]:
if not model:
return None
if self._compiled is None:
self._compiled = self._compile()
for pattern, model_info in self._compiled:
if pattern.search(model) is not None:
return dict(model_info)
return None
return next(
(rule.provider for rule in self.routing_rules if rule.pattern.search(model) is not None),
None,
)
def match_capabilities(self, model: str) -> Optional[dict]:
if not model:
return None
matched = tuple(rule.model_info for rule in self.capability_rules if rule.pattern.search(model) is not None)
if not matched:
return None
return {key: value for model_info in matched for key, value in model_info.items()}
_registry = _FallbackGeneralizations()
def set_fallback_generalizations(rules: Optional[list[dict]]) -> None:
"""Install the active rule list and invalidate the compiled-regex cache.
def set_fallback_generalizations(rules: Optional[list]) -> None:
"""Install the active rule list, compiling and classifying each rule.
``extends`` inheritance is resolved here, once, before the rules are stored.
Called once when the model cost map is loaded (and again on any reload).
Legacy ``extends`` inheritance is resolved here, once, before classification;
a legacy rule mixing ``litellm_provider`` with capability keys installs as both
kinds. Malformed and invalid-regex rules are warned about and skipped. Called
once when the model cost map is loaded (and again on any reload).
"""
_registry.set_rules(_resolve_extends(rules) if isinstance(rules, list) else rules)
_registry.set_rules(rules)
def get_fallback_generalization_rules() -> list[dict]:
def get_fallback_generalization_rules() -> list:
"""Return the raw rule list (read-only view for callers/tests)."""
return _registry.rules
def match_fallback_generalization(model: str) -> Optional[dict]:
"""Return the ``model_info`` of the first rule whose regex matches ``model``.
def match_routing_generalization(model: str) -> Optional[str]:
"""Return the provider of the first routing rule whose regex matches ``model``.
O(number of rules). Only call this once exact lookups have missed.
"""
return _registry.match(model)
return _registry.match_routing(model)
def match_capability_generalizations(model: str) -> Optional[dict]:
"""Return the union of the ``model_info`` of every capability rule matching ``model``.
Later rules override earlier ones on key conflicts. Returns ``None`` when no
capability rule matches. O(number of rules); only call once exact lookups have missed.
"""
return _registry.match_capabilities(model)

View file

@ -2,26 +2,8 @@ from typing import Optional
from litellm.llms.openai.data_residency import infer_openai_data_residency
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
OPTIONAL_KWARGS_KEYS = frozenset(
AWS_CREDENTIAL_KWARGS_KEYS = frozenset(
{
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_username",
"azure_password",
"azure_scope",
"timeout",
"gcs_bucket_name",
"bucket_name",
"vertex_credentials",
"vertex_project",
"vertex_location",
"vertex_ai_project",
"vertex_ai_location",
"vertex_ai_credentials",
"aws_region_name",
"aws_access_key_id",
"aws_secret_access_key",
@ -34,14 +16,40 @@ OPTIONAL_KWARGS_KEYS = frozenset(
"aws_external_id",
"aws_bedrock_runtime_endpoint",
"aws_bedrock_project_id",
"tpm",
"rpm",
"itpm",
"otpm",
"use_xai_oauth",
}
)
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
OPTIONAL_KWARGS_KEYS = (
frozenset(
{
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_username",
"azure_password",
"azure_scope",
"timeout",
"gcs_bucket_name",
"bucket_name",
"vertex_credentials",
"vertex_project",
"vertex_location",
"vertex_ai_project",
"vertex_ai_location",
"vertex_ai_credentials",
"tpm",
"rpm",
"itpm",
"otpm",
"use_xai_oauth",
}
)
| AWS_CREDENTIAL_KWARGS_KEYS
)
# Backward-compatible alias for existing imports/tests.
_OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS

View file

@ -4,7 +4,7 @@ from urllib.parse import urlparse
import litellm
from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH
from litellm.litellm_core_utils.fallback_generalizations import (
match_fallback_generalization,
match_routing_generalization,
)
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.secret_managers.main import get_secret, get_secret_str
@ -474,12 +474,10 @@ def get_llm_provider(
custom_llm_provider = "sap"
# Last resort for an otherwise-unknown model: a declarative
# fallback-generalization rule (e.g. routes future claude-* to anthropic).
# fallback-generalization routing rule (e.g. routes future claude-* to anthropic).
# Exact provider matches above always win; this only runs on a miss.
if not custom_llm_provider:
generalization = match_fallback_generalization(model)
if generalization is not None:
custom_llm_provider = generalization.get("litellm_provider") or None
custom_llm_provider = match_routing_generalization(model)
if not custom_llm_provider:
if litellm.suppress_debug_info is False:

View file

@ -38,6 +38,7 @@ from litellm import (
)
from litellm._logging import _is_debugging_on, _redact_string, verbose_logger
from litellm.exceptions import (
BudgetExceededError,
validate_rate_limit_category,
validate_rate_limit_type,
)
@ -73,6 +74,7 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
from litellm.litellm_core_utils.redact_messages import (
redact_message_input_output_from_custom_logger,
redact_message_input_output_from_logging,
redact_streaming_responses_for_custom_logger,
)
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.llms.base_llm.search.transformation import SearchResponse
@ -924,7 +926,6 @@ class Logging(LiteLLMLoggingBaseClass):
def pre_call(self, input, api_key, model=None, additional_args={}):
# Log the exact input to the LLM API
litellm.error_logs["PRE_CALL"] = locals()
try:
self._pre_call(
input=input,
@ -1134,7 +1135,6 @@ class Logging(LiteLLMLoggingBaseClass):
def post_call(self, original_response, input=None, api_key=None, additional_args={}):
# Log the exact result from the LLM API, for streaming - log the type of response received
litellm.error_logs["POST_CALL"] = locals()
if isinstance(original_response, dict):
original_response = json.dumps(original_response, default=str)
try:
@ -1531,6 +1531,9 @@ class Logging(LiteLLMLoggingBaseClass):
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
and litellm_params.get(CallTypes.atranscription.value, False) is not True
and litellm_params.get(CallTypes.allm_passthrough_route.value, False) is not True
and litellm_params.get(CallTypes.aanthropic_messages.value, False) is not True
and litellm_params.get(CallTypes.agenerate_content.value, False) is not True
and litellm_params.get(CallTypes.agenerate_content_stream.value, False) is not True
)
def _is_assembled_stream_success(self, result=None) -> bool:
@ -2576,6 +2579,9 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details = callback.redact_standard_logging_payload_from_model_call_details(
model_call_details=model_call_details
)
model_call_details = redact_streaming_responses_for_custom_logger(
model_call_details=model_call_details, custom_logger=callback
)
##################################
if self.stream is True:
if "async_complete_streaming_response" in model_call_details:
@ -3070,7 +3076,7 @@ class Logging(LiteLLMLoggingBaseClass):
def get_combined_callback_list(self, dynamic_success_callbacks: Optional[List], global_callbacks: List) -> List:
if dynamic_success_callbacks is None:
return list(global_callbacks)
return list(set(dynamic_success_callbacks + global_callbacks))
return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks))
def _remove_internal_litellm_callbacks(self, callbacks: List) -> List:
"""
@ -4595,6 +4601,10 @@ class StandardLoggingPayloadSetup:
user_api_key_spend=None,
user_api_key_max_budget=None,
user_api_key_budget_reset_at=None,
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_org_alias=None,
@ -4941,6 +4951,7 @@ class StandardLoggingPayloadSetup:
rate_limit_category = validate_rate_limit_category(getattr(original_exception, "category", None))
rate_limit_type = validate_rate_limit_type(getattr(original_exception, "rate_limit_type", None))
budget_error = original_exception if isinstance(original_exception, BudgetExceededError) else None
return StandardLoggingPayloadErrorInformation(
error_code=error_status,
@ -4950,6 +4961,10 @@ class StandardLoggingPayloadSetup:
error_message=error_message,
error_rate_limit_category=rate_limit_category,
error_rate_limit_type=rate_limit_type,
error_budget_entity_type=budget_error.entity_type if budget_error else None,
error_budget_entity_id=budget_error.entity_id if budget_error else None,
error_budget_limit=budget_error.max_budget if budget_error else None,
error_budget_spend=budget_error.current_cost if budget_error else None,
)
@staticmethod
@ -5208,10 +5223,15 @@ def get_standard_logging_object_payload(
call_type = kwargs.get("call_type")
cache_hit = kwargs.get("cache_hit", False)
# Extract usage as a plain dict, avoiding Pydantic round-trip
usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict(
raw_usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj=response_obj,
combined_usage_object=cast(Optional[Usage], kwargs.get("combined_usage_object")),
)
usage_dict = (
{**raw_usage_dict, "output_image_count": len(init_response_obj.data)}
if isinstance(init_response_obj, ImageResponse) and init_response_obj.data
else raw_usage_dict
)
id = response_obj.get("id", kwargs.get("litellm_call_id"))
@ -5421,6 +5441,10 @@ def get_standard_logging_metadata(
user_api_key_spend=None,
user_api_key_max_budget=None,
user_api_key_budget_reset_at=None,
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_org_alias=None,
@ -5520,6 +5544,10 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
user_api_key_team_id=str("test_team"),
user_api_key_user_id=str("test_user"),
user_api_key_team_alias=str("test_team_alias"),
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_org_id=None,
spend_logs_metadata=None,
requester_ip_address=str("127.0.0.1"),

View file

@ -0,0 +1,139 @@
"""
Provider-neutral graduated tiered pricing calculation.
Shared by provider cost calculators (e.g. Dashscope) and the proxy budget
reservation logic so neither has to depend on the other.
"""
from typing import List, Optional, Union
def _coerce_cost_per_token(value: Union[float, int, str, None]) -> float:
"""
Coerce a per-token cost into a float.
Model cost values loaded from YAML config may arrive as strings (e.g.
scientific notation like "4e-07"), which would break arithmetic.
"""
if value is None:
return 0.0
if isinstance(value, str):
try:
return float(value)
except ValueError:
return 0.0
return float(value)
def calculate_tiered_cost(
tokens: int,
tiered_pricing: List[dict],
cost_key: str,
fallback_cost_key: Optional[str] = None,
) -> float:
"""
Calculate cost for a given number of tokens based on a true tiered pricing structure.
This function iterates through sorted pricing tiers, calculates the cost for the
number of tokens that fall into each tier's range, and sums them up to get the total cost.
Args:
tokens (int): The total number of tokens to calculate the cost for.
tiered_pricing (List[dict]): A list of dictionaries, where each dictionary
represents a pricing tier.
cost_key (str): The key in the tier dictionary that holds the per-token cost
(e.g., 'input_cost_per_token').
fallback_cost_key (Optional[str], optional): A fallback key to use if the
primary `cost_key` is not found in a tier. Defaults to None.
Returns:
float: The total calculated cost for the given tokens.
Example:
>>> tiered_pricing = [
... {"range": [0, 100000], "input_cost_per_token": 0.0001},
... {"range": [100000, 500000], "input_cost_per_token": 0.00005},
... ]
Calculating cost for 150,000 tokens:
(100,000 * 0.0001) + (50,000 * 0.00005) = $12.5
"""
if not tiered_pricing or tokens <= 0:
return 0.0
total_cost = 0.0
tokens_processed = 0
sorted_tiers = sorted(tiered_pricing, key=lambda x: x.get("range", [0, 0])[0])
for tier in sorted_tiers:
if tokens_processed >= tokens:
break
tier_range = tier.get("range", [])
if len(tier_range) != 2:
continue
range_start, range_end = tier_range
if tokens <= range_start:
continue
tier_start = max(range_start, tokens_processed)
tier_end = min(range_end, tokens)
if tier_end > tier_start:
tokens_in_tier = tier_end - tier_start
cost_per_token = tier.get(cost_key) or tier.get(fallback_cost_key, 0)
total_cost += tokens_in_tier * _coerce_cost_per_token(cost_per_token)
tokens_processed = tier_end
# After loop, check if any tokens remain (i.e., tokens > highest tier's end range)
# and charge them at the last tier's rate.
if tokens_processed < tokens and sorted_tiers:
last_tier = sorted_tiers[-1]
remaining_tokens = tokens - tokens_processed
cost_per_token = last_tier.get(cost_key) or last_tier.get(fallback_cost_key, 0)
total_cost += remaining_tokens * _coerce_cost_per_token(cost_per_token)
return total_cost
def select_tier_for_input(
tiered_pricing: List[dict],
input_tokens: int,
) -> Optional[dict]:
"""
Select the pricing tier for a request based on its total input token count.
Alibaba Model Studio (Dashscope) tiered pricing is all-or-nothing: the tier is
chosen by the total input tokens of a single request and every token in the
request (input and output) is billed at that one tier's rate, rather than
graduated income-tax-style slicing. A tier matches when
``range_start < input_tokens <= range_end`` (so a request of exactly
``range_end`` tokens stays in the lower tier, matching the official
``0 < Token <= 32K`` phrasing). Requests above the highest declared range fall
back to the last (most expensive) tier.
"""
if not tiered_pricing or input_tokens <= 0:
return None
sorted_tiers = sorted(tiered_pricing, key=lambda t: t.get("range", [0, 0])[0])
valid_tiers = [tier for tier in sorted_tiers if len(tier.get("range", [])) == 2]
if not valid_tiers:
return None
matching = [tier for tier in valid_tiers if tier["range"][0] < input_tokens <= tier["range"][1]]
if matching:
return matching[0]
return valid_tiers[-1]
def tier_rate(
tier: dict,
cost_key: str,
fallback_cost_key: Optional[str] = None,
) -> float:
"""Read a per-token rate from a tier, coercing YAML string costs to float."""
raw = tier.get(cost_key) or tier.get(fallback_cost_key, 0)
return _coerce_cost_per_token(raw)

View file

@ -445,6 +445,7 @@ class PromptTokensDetailsResult(TypedDict):
text_tokens: int
audio_tokens: int
image_tokens: int
video_tokens: int
character_count: int
image_count: int
video_length_seconds: float
@ -473,6 +474,7 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
)
audio_tokens = cast(Optional[int], getattr(usage.prompt_tokens_details, "audio_tokens", 0)) or 0
image_tokens = cast(Optional[int], getattr(usage.prompt_tokens_details, "image_tokens", 0)) or 0
video_tokens = _coerce_token_count(getattr(usage.prompt_tokens_details, "video_tokens", 0))
character_count = (
cast(
Optional[int],
@ -503,6 +505,7 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
text_tokens=text_tokens,
audio_tokens=audio_tokens,
image_tokens=image_tokens,
video_tokens=video_tokens,
character_count=character_count,
image_count=image_count,
video_length_seconds=float(video_length_seconds),
@ -515,6 +518,7 @@ class CompletionTokensDetailsResult(TypedDict):
text_tokens: int
reasoning_tokens: int
image_tokens: int
video_tokens: int
def _parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsResult:
@ -546,12 +550,14 @@ def _parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsRes
)
or 0
)
video_tokens = _coerce_token_count(getattr(usage.completion_tokens_details, "video_tokens", 0))
return CompletionTokensDetailsResult(
audio_tokens=audio_tokens,
text_tokens=text_tokens,
reasoning_tokens=reasoning_tokens,
image_tokens=image_tokens,
video_tokens=video_tokens,
)
@ -586,6 +592,13 @@ def _calculate_input_cost(
image_token_cost_key = "input_cost_per_token"
prompt_cost += calculate_cost_component(model_info, image_token_cost_key, prompt_tokens_details["image_tokens"])
### VIDEO TOKEN COST
if prompt_tokens_details["video_tokens"]:
video_token_cost_key = "input_cost_per_video_token"
if model_info.get(video_token_cost_key) is None:
video_token_cost_key = "input_cost_per_token"
prompt_cost += calculate_cost_component(model_info, video_token_cost_key, prompt_tokens_details["video_tokens"])
### CACHE WRITING COST - Now uses tiered pricing
if (
prompt_tokens_details["cache_creation_tokens"]
@ -698,6 +711,7 @@ def generic_cost_per_token(
text_tokens=usage.prompt_tokens,
audio_tokens=0,
image_tokens=0,
video_tokens=0,
character_count=0,
image_count=0,
video_length_seconds=0.0,
@ -716,13 +730,14 @@ def generic_cost_per_token(
audio_tokens = prompt_tokens_details["audio_tokens"]
cache_creation = prompt_tokens_details["cache_creation_tokens"]
image_tokens = prompt_tokens_details["image_tokens"]
video_tokens = prompt_tokens_details["video_tokens"]
# Check for double-counting: sum of details > prompt_tokens means overlap
total_details = text_tokens + cache_hit + audio_tokens + cache_creation + image_tokens
total_details = text_tokens + cache_hit + audio_tokens + cache_creation + image_tokens + video_tokens
has_double_counting = cache_hit > 0 and total_details > usage.prompt_tokens
if (text_tokens == 0 and prompt_tokens_details["image_count"] == 0) or has_double_counting:
text_tokens = usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens
text_tokens = usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens
# Clamp to zero: inconsistent streaming usage
if text_tokens < 0:
text_tokens = 0
@ -751,6 +766,7 @@ def generic_cost_per_token(
audio_tokens = 0
reasoning_tokens = 0
image_tokens = 0
video_tokens = 0
is_text_tokens_total = False
if usage.completion_tokens_details is not None:
completion_tokens_details = _parse_completion_tokens_details(usage)
@ -758,19 +774,20 @@ def generic_cost_per_token(
text_tokens = completion_tokens_details["text_tokens"]
reasoning_tokens = completion_tokens_details["reasoning_tokens"]
image_tokens = completion_tokens_details["image_tokens"]
video_tokens = completion_tokens_details["video_tokens"]
# Handle text_tokens calculation:
# 1. If text_tokens is explicitly provided and > 0, use it
# 2. If there's a breakdown (reasoning/audio/image tokens), calculate text_tokens as the remainder
# 2. If there's a breakdown (reasoning/audio/image/video tokens), calculate text_tokens as the remainder
# 3. If no breakdown at all, assume all completion_tokens are text_tokens
has_token_breakdown = image_tokens > 0 or audio_tokens > 0 or reasoning_tokens > 0
has_token_breakdown = image_tokens > 0 or audio_tokens > 0 or reasoning_tokens > 0 or video_tokens > 0
if text_tokens == 0:
if has_token_breakdown:
# Calculate text tokens as remainder when we have a breakdown
# This handles cases like OpenAI's reasoning models where text_tokens isn't provided
text_tokens = max(
0,
usage.completion_tokens - reasoning_tokens - audio_tokens - image_tokens,
usage.completion_tokens - reasoning_tokens - audio_tokens - image_tokens - video_tokens,
)
else:
# No breakdown at all, all tokens are text tokens
@ -803,6 +820,14 @@ def generic_cost_per_token(
)
completion_cost += float(image_tokens) * _output_cost_per_image_token
## VIDEO COST
if not is_text_tokens_total and video_tokens and video_tokens > 0:
_output_cost_per_video_token = _get_cost_per_unit(model_info, "output_cost_per_video_token", None)
_output_cost_per_video_token = (
_output_cost_per_video_token if _output_cost_per_video_token is not None else completion_base_cost
)
completion_cost += float(video_tokens) * _output_cost_per_video_token
## REGIONAL DATA-RESIDENCY UPLIFT
# Applied as a flat multiplier across all token costs for the request
# when the upstream is a regionalized OpenAI host (eu./us.api.openai.com).

View file

@ -5494,3 +5494,56 @@ def has_tool_with_name(tools: Any, tool_name: str) -> bool:
elif tool.get("name") == tool_name:
return True
return False
def resolve_structured_messages(
messages: list[dict[str, Any]] | None,
request_kwargs: dict[str, Any],
) -> list[dict[str, Any]] | None:
"""
Normalize a request's messages to OpenAI-spec chat-completions shape,
regardless of which API surface produced them (chat completions,
Anthropic /v1/messages, Responses API ``input``, etc).
Returns ``messages`` unchanged if already present. Otherwise dispatches
through the guardrail translation handlers (the same per-surface
conversion logic guardrails use) to convert e.g. Responses API ``input``
into a message list. Returns ``None`` if no messages could be resolved.
"""
if messages:
return messages
from litellm.litellm_core_utils.api_route_to_call_types import (
get_call_types_for_route,
)
from litellm.llms import load_guardrail_translation_mappings
from litellm.types.utils import CallTypes
mappings = load_guardrail_translation_mappings()
call_type: CallTypes | None = None
# 1. Try route-based inference from proxy metadata
route = request_kwargs.get("litellm_metadata", {}).get("user_api_key_request_route")
if route:
call_types_list = get_call_types_for_route(route)
if call_types_list:
for ct in call_types_list:
if ct in mappings:
call_type = ct
break
# 2. Fallback: try each mapped handler until one produces messages
handlers_to_try: list[Any] = []
if call_type is not None and call_type in mappings:
handlers_to_try.append(mappings[call_type]())
else:
handlers_to_try.extend(handler_cls() for handler_cls in mappings.values())
for handler in handlers_to_try:
structured = handler.get_structured_messages(request_kwargs)
if structured:
return [
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
for msg in structured
]
return None

View file

@ -38,10 +38,61 @@ def redact_message_input_output_from_custom_logger(
litellm_logging_obj: LiteLLMLoggingObject, result, custom_logger: CustomLogger
):
if hasattr(custom_logger, "message_logging") and custom_logger.message_logging is not True:
return perform_redaction(litellm_logging_obj.model_call_details, result)
return perform_redaction(litellm_logging_obj.model_call_details, result, redact_streaming_responses=False)
return result
def redact_streaming_responses_for_custom_logger(model_call_details: dict, custom_logger: CustomLogger) -> dict:
"""
Returns a copy of model_call_details whose streaming response entries are redacted deepcopies
when the custom logger has opted out of message logging. The shared model_call_details is left
untouched so other callbacks still receive the unredacted response.
"""
if not (hasattr(custom_logger, "message_logging") and custom_logger.message_logging is not True):
return model_call_details
redacted_entries = {
streaming_key: _redacted_streaming_response_copy(model_call_details[streaming_key])
for streaming_key in ("complete_streaming_response", "async_complete_streaming_response")
if model_call_details.get(streaming_key) is not None
}
if not redacted_entries:
return model_call_details
return {**model_call_details, **redacted_entries}
def _redacted_streaming_response_copy(streaming_response):
redacted_response = copy.deepcopy(streaming_response)
_redact_streaming_response(redacted_response)
return redacted_response
def _redact_streaming_response(streaming_response):
if hasattr(streaming_response, "choices"):
for choice in streaming_response.choices:
_redact_choice_content(choice)
redact_vertex_ai_metadata_from_logged_object(streaming_response)
elif hasattr(streaming_response, "output"):
_redact_responses_api_output(streaming_response.output)
if hasattr(streaming_response, "reasoning") and streaming_response.reasoning is not None:
streaming_response.reasoning = None
def _redact_tool_calls(tool_calls) -> None:
"""Redact tool call arguments (assistant tool calls carry prompt-derived data)."""
if not tool_calls:
return
for tool_call in tool_calls:
function = getattr(tool_call, "function", None)
if function is not None and hasattr(function, "arguments"):
function.arguments = "redacted-by-litellm"
def _redact_function_call(function_call) -> None:
"""Redact legacy assistant function_call arguments."""
if function_call is not None and hasattr(function_call, "arguments"):
function_call.arguments = "redacted-by-litellm"
def _redact_choice_content(choice):
"""Helper to redact content in a choice (message or delta)."""
if isinstance(choice, litellm.Choices):
@ -50,12 +101,16 @@ def _redact_choice_content(choice):
choice.message.reasoning_content = "redacted-by-litellm"
if hasattr(choice.message, "thinking_blocks"):
choice.message.thinking_blocks = None
_redact_tool_calls(getattr(choice.message, "tool_calls", None))
_redact_function_call(getattr(choice.message, "function_call", None))
elif isinstance(choice, litellm.utils.StreamingChoices):
choice.delta.content = "redacted-by-litellm"
if hasattr(choice.delta, "reasoning_content"):
choice.delta.reasoning_content = "redacted-by-litellm"
if hasattr(choice.delta, "thinking_blocks"):
choice.delta.thinking_blocks = None
_redact_tool_calls(getattr(choice.delta, "tool_calls", None))
_redact_function_call(getattr(choice.delta, "function_call", None))
def _redact_responses_api_output(output_items):
@ -76,6 +131,9 @@ def _redact_responses_api_output(output_items):
if hasattr(summary_item, "text"):
summary_item.text = "redacted-by-litellm"
if hasattr(output_item, "type") and output_item.type == "function_call" and hasattr(output_item, "arguments"):
output_item.arguments = "redacted-by-litellm"
def _redact_responses_api_output_dict(output_items, redacted_str: str):
"""Helper to redact ResponsesAPIResponse output items in dict form."""
@ -96,6 +154,9 @@ def _redact_responses_api_output_dict(output_items, redacted_str: str):
if isinstance(summary_item, dict) and "text" in summary_item:
summary_item["text"] = redacted_str
if output_item.get("type") == "function_call" and "arguments" in output_item:
output_item["arguments"] = redacted_str
def _redact_standard_logging_object(model_call_details: dict):
"""Redact messages and response inside standard_logging_object if present."""
@ -127,6 +188,19 @@ def _redact_standard_logging_object(model_call_details: dict):
standard_logging_object["response"] = {"text": redacted_str}
def _redact_tool_calls_dict(message: dict, redacted_str: str) -> None:
"""Redact tool call / function_call arguments in a dict-form message or delta."""
tool_calls = message.get("tool_calls")
if isinstance(tool_calls, list):
for tool_call in tool_calls:
if isinstance(tool_call, dict) and isinstance(tool_call.get("function"), dict):
tool_call["function"]["arguments"] = redacted_str
function_call = message.get("function_call")
if isinstance(function_call, dict) and "arguments" in function_call:
function_call["arguments"] = redacted_str
def _redact_model_response_dict_choices(choices, redacted_str: str):
for choice in choices:
if isinstance(choice, dict):
@ -138,6 +212,7 @@ def _redact_model_response_dict_choices(choices, redacted_str: str):
choice["message"]["thinking_blocks"] = None
if "audio" in choice["message"]:
choice["message"]["audio"] = None
_redact_tool_calls_dict(choice["message"], redacted_str)
elif "delta" in choice and isinstance(choice["delta"], dict):
choice["delta"]["content"] = redacted_str
if "reasoning_content" in choice["delta"]:
@ -146,13 +221,18 @@ def _redact_model_response_dict_choices(choices, redacted_str: str):
choice["delta"]["thinking_blocks"] = None
if "audio" in choice["delta"]:
choice["delta"]["audio"] = None
_redact_tool_calls_dict(choice["delta"], redacted_str)
else:
_redact_choice_content(choice)
def perform_redaction(model_call_details: dict, result):
def perform_redaction(model_call_details: dict, result, redact_streaming_responses: bool = True):
"""
Performs the actual redaction on the logging object and result.
redact_streaming_responses=False skips the in-place redaction of the shared streaming
response entries; per-callback redaction hands each opted-out callback its own redacted
copy via redact_streaming_responses_for_custom_logger instead.
"""
# Redact model_call_details
model_call_details["messages"] = [{"role": "user", "content": "redacted-by-litellm"}]
@ -162,17 +242,9 @@ def perform_redaction(model_call_details: dict, result):
redact_vertex_ai_metadata_from_litellm_params(model_call_details)
# Redact streaming response
if model_call_details.get("stream", False) is True and "complete_streaming_response" in model_call_details:
_streaming_response = model_call_details["complete_streaming_response"]
if hasattr(_streaming_response, "choices"):
for choice in _streaming_response.choices:
_redact_choice_content(choice)
redact_vertex_ai_metadata_from_logged_object(_streaming_response)
elif hasattr(_streaming_response, "output"):
_redact_responses_api_output(_streaming_response.output)
# Redact reasoning field in ResponsesAPIResponse
if hasattr(_streaming_response, "reasoning") and _streaming_response.reasoning is not None:
_streaming_response.reasoning = None
if redact_streaming_responses and model_call_details.get("stream", False) is True:
for _streaming_key in ("complete_streaming_response", "async_complete_streaming_response"):
_redact_streaming_response(model_call_details.get(_streaming_key))
# Redact result
if result is not None:

View file

@ -9,6 +9,8 @@ secrets from strings without depending on the logging-configuration module.
import re
from typing import List
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH
_REDACTED = "REDACTED"
@ -30,7 +32,7 @@ def _build_secret_patterns() -> "re.Pattern[str]":
# Basic auth headers
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
# OpenAI / Anthropic sk- prefixed keys
r"sk-[A-Za-z0-9\-_]{20,}",
rf"sk-[A-Za-z0-9\-_]{{{MINIMUM_CUSTOM_KEY_LENGTH - len('sk-')},}}",
# Generic api_key / api-key / apikey (handles 'key': 'value' dict repr)
r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}",
# x-api-key / api-key header values (handles 'key': 'value' dict repr)

View file

@ -467,6 +467,7 @@ class ChunkProcessor:
cache_read_input_tokens: Optional[int] = None
completion_tokens_details: Optional[CompletionTokensDetails] = None
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
cost: Optional[float] = None
if "prompt_tokens" in usage_chunk:
prompt_tokens = usage_chunk.get("prompt_tokens", 0) or 0
@ -476,6 +477,8 @@ class ChunkProcessor:
cache_creation_input_tokens = usage_chunk.get("cache_creation_input_tokens")
if "cache_read_input_tokens" in usage_chunk:
cache_read_input_tokens = usage_chunk.get("cache_read_input_tokens")
if "cost" in usage_chunk:
cost = usage_chunk.get("cost")
if hasattr(usage_chunk, "completion_tokens_details"):
if isinstance(usage_chunk.completion_tokens_details, dict):
completion_tokens_details = CompletionTokensDetails(**usage_chunk.completion_tokens_details)
@ -494,6 +497,7 @@ class ChunkProcessor:
"cache_read_input_tokens": cache_read_input_tokens,
"completion_tokens_details": completion_tokens_details,
"prompt_tokens_details": prompt_tokens_details,
"cost": cost,
}
def count_reasoning_tokens(self, response: ModelResponse) -> Optional[int]:
@ -512,6 +516,22 @@ class ChunkProcessor:
return reasoning_tokens
@staticmethod
def _extract_usage_chunk(chunk: dict[str, Any] | ModelResponse | ModelResponseStream) -> Usage | None:
usage_chunk: Usage | dict[str, Any] | None = None
if hasattr(chunk, "usage") and chunk.usage is not None:
usage_chunk = chunk.usage
elif "usage" in chunk:
usage_chunk = chunk["usage"]
elif (isinstance(chunk, ModelResponse) or isinstance(chunk, ModelResponseStream)) and hasattr(
chunk, "_hidden_params"
):
usage_chunk = chunk._hidden_params.get("usage", None)
if isinstance(usage_chunk, dict):
return Usage(**usage_chunk)
return usage_chunk
def _calculate_usage_per_chunk(
self,
chunks: List[Union[Dict[str, Any], ModelResponse]],
@ -548,18 +568,12 @@ class ChunkProcessor:
# is last-wins, so without preserving this separately the 1h breakdown is
# lost and 1h cache writes get billed at the 5m rate.
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
cost: Optional[float] = None
for chunk in chunks:
usage_chunk: Optional[Usage] = None
if "usage" in chunk:
usage_chunk = chunk["usage"]
elif (isinstance(chunk, ModelResponse) or isinstance(chunk, ModelResponseStream)) and hasattr(
chunk, "_hidden_params"
):
usage_chunk = chunk._hidden_params.get("usage", None)
usage_chunk = self._extract_usage_chunk(chunk)
if usage_chunk is not None:
if isinstance(usage_chunk, dict):
usage_chunk = Usage(**usage_chunk)
usage_chunk_dict = self._usage_chunk_calculation_helper(usage_chunk)
if usage_chunk_dict["prompt_tokens"] is not None and usage_chunk_dict["prompt_tokens"] > 0:
prompt_tokens = usage_chunk_dict["prompt_tokens"]
@ -610,6 +624,9 @@ class ChunkProcessor:
prompt_tokens_details, cache_creation_token_details
)
if usage_chunk_dict["cost"] is not None:
cost = usage_chunk_dict["cost"]
prompt_tokens_details = self._attach_cache_creation_token_details(
prompt_tokens_details, cache_creation_token_details
)
@ -629,6 +646,7 @@ class ChunkProcessor:
web_search_requests=web_search_requests,
completion_tokens_details=completion_tokens_details,
prompt_tokens_details=prompt_tokens_details,
cost=cost,
)
@staticmethod
@ -727,6 +745,7 @@ class ChunkProcessor:
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = calculated_usage_per_chunk[
"prompt_tokens_details"
]
cost: Optional[float] = calculated_usage_per_chunk["cost"]
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(model=model, messages=messages)
@ -784,6 +803,9 @@ class ChunkProcessor:
else:
returned_usage.prompt_tokens_details.web_search_requests = web_search_requests
if cost is not None:
setattr(returned_usage, "cost", cost)
# Return a new usage object with the new values
returned_usage = Usage(**returned_usage.model_dump())

View file

@ -962,10 +962,11 @@ class CustomStreamWrapper:
if self.custom_llm_provider == "bedrock" and "trace" in model_response:
return model_response
# Default - return StopIteration
if hasattr(model_response, "usage"):
self.chunks.append(model_response)
raise StopIteration
# Don't raise StopIteration here - some providers (like OpenRouter)
# send usage/cost data in chunks after the finish_reason chunk
if hasattr(model_response, "usage") and model_response.usage is not None:
return model_response
return
# flush any remaining holding chunk
if len(self.holding_chunk) > 0:
if model_response.choices[0].delta.content is None:
@ -1474,12 +1475,16 @@ class CustomStreamWrapper:
self.tool_call = True
if hasattr(chunk, "usage") and chunk.usage is not None:
model_response.usage = chunk.usage
## RETURN ARG
return self.return_processed_chunk_logic(
result = self.return_processed_chunk_logic(
completion_obj=completion_obj,
model_response=model_response, # type: ignore
response_obj=response_obj,
)
return result
except StopIteration:
raise StopIteration
@ -1686,6 +1691,21 @@ class CustomStreamWrapper:
model_response.choices[0].finish_reason = "tool_calls"
return model_response
@staticmethod
def _propagate_usage_cost_to_hidden_params(
response: "ModelResponse",
) -> None:
"""
If the assembled response carries a provider-reported cost on
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
_usage = getattr(response, "usage", None)
if _usage is not None and hasattr(_usage, "cost") and _usage.cost is not None:
if "additional_headers" not in response._hidden_params:
response._hidden_params["additional_headers"] = {}
response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = float(_usage.cost)
def __next__(self) -> "ModelResponseStream":
cache_hit = False
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
@ -1741,6 +1761,10 @@ class CustomStreamWrapper:
# hasattr(response, "usage") is always True — must check
# `is not None` to avoid running this path on every chunk.
if getattr(response, "usage", None) is not None:
usage_to_preserve = response.usage
if usage_to_preserve:
response._hidden_params["usage"] = usage_to_preserve
obj_dict = response.model_dump()
if "usage" in obj_dict:
@ -1789,6 +1813,8 @@ class CustomStreamWrapper:
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
setattr(
response,
"usage",
@ -1974,97 +2000,7 @@ class CustomStreamWrapper:
self.chunks.append(processed_chunk)
return processed_chunk
except (StopAsyncIteration, StopIteration):
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
setattr(
response,
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(
self.async_cache_streaming_response(
processed_chunk=_copy,
cache_hit=cache_hit,
)
)
# Update hidden_params with final usage from
# stream_chunk_builder (see sync __next__ for full comment).
if (
self.stream_options is None
and complete_streaming_response is not None
and self._last_returned_hidden_params is not None
):
final_usage = getattr(complete_streaming_response, "usage", None)
if final_usage is not None:
self._last_returned_hidden_params["usage"] = final_usage
if self.sent_stream_usage is False and self.send_stream_usage is True:
self.sent_stream_usage = True
return response
_deferred_cb = getattr(
self.logging_obj,
"_on_deferred_stream_complete",
None,
)
if _deferred_cb is not None:
# Proxy has post-call guardrails. Store the assembled
# response so the outer streaming consumer
# (ProxyLogging.async_post_call_streaming_iterator_hook)
# can fire the deferred callback AFTER all guardrail
# end-of-stream blocks complete. Scheduling here via
# create_task would race with unified_guardrail's
# end-of-stream block for short-stream providers.
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
complete_streaming_response,
cache_hit,
)
else:
# prefer_async_handlers routes CustomLogger to async_success_handler
# when consumers use ``async for`` on sync-SDK streams. Legacy string
# callbacks still run via executor.submit inside dispatch_success_handlers.
asyncio.create_task(
self.logging_obj.dispatch_success_handlers(
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
)
raise StopAsyncIteration # Re-raise StopIteration
else:
self.sent_last_chunk = True
processed_chunk = self.finish_reason_handler()
return processed_chunk
return await self._finalize_completed_stream(cache_hit=cache_hit)
except httpx.TimeoutException as e: # if httpx read timeout error occues
traceback_exception = traceback.format_exc()
## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT
@ -2079,20 +2015,122 @@ class CustomStreamWrapper:
# Handle any exceptions that might occur during streaming
asyncio.create_task(self.logging_obj.async_failure_handler(e, traceback_exception))
self._handle_stream_fallback_error(e)
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
if self.received_finish_reason is None:
self._log_stream_failure_and_raise(e)
return await self._finalize_completed_stream(cache_hit=cache_hit)
except Exception as e:
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:
self._record_partial_usage_for_failure()
## LOGGING
threading.Thread(
target=self.logging_obj.failure_handler,
args=(e, traceback_exception),
).start() # log response
# Handle any exceptions that might occur during streaming
asyncio.create_task(
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
self._log_stream_failure_and_raise(e)
async def _finalize_completed_stream(self, cache_hit: bool) -> "ModelResponseStream":
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
self._handle_stream_fallback_error(e)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
setattr(
response,
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(
self.async_cache_streaming_response(
processed_chunk=_copy,
cache_hit=cache_hit,
)
)
# Update hidden_params with final usage from
# stream_chunk_builder (see sync __next__ for full comment).
if (
self.stream_options is None
and complete_streaming_response is not None
and self._last_returned_hidden_params is not None
):
final_usage = getattr(complete_streaming_response, "usage", None)
if final_usage is not None:
self._last_returned_hidden_params["usage"] = final_usage
if self.sent_stream_usage is False and self.send_stream_usage is True:
self.sent_stream_usage = True
return response
_deferred_cb = getattr(
self.logging_obj,
"_on_deferred_stream_complete",
None,
)
if _deferred_cb is not None:
# Proxy has post-call guardrails. Store the assembled
# response so the outer streaming consumer
# (ProxyLogging.async_post_call_streaming_iterator_hook)
# can fire the deferred callback AFTER all guardrail
# end-of-stream blocks complete. Scheduling here via
# create_task would race with unified_guardrail's
# end-of-stream block for short-stream providers.
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
complete_streaming_response,
cache_hit,
)
else:
# prefer_async_handlers routes CustomLogger to async_success_handler
# when consumers use ``async for`` on sync-SDK streams. Legacy string
# callbacks still run via executor.submit inside dispatch_success_handlers.
asyncio.create_task(
self.logging_obj.dispatch_success_handlers(
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
)
raise StopAsyncIteration # Re-raise StopIteration
else:
self.sent_last_chunk = True
processed_chunk = self.finish_reason_handler()
return processed_chunk
def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn:
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:
self._record_partial_usage_for_failure()
## LOGGING
threading.Thread(
target=self.logging_obj.failure_handler,
args=(e, traceback_exception),
).start() # log response
# Handle any exceptions that might occur during streaming
asyncio.create_task(
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
)
self._handle_stream_fallback_error(e)
def _record_partial_usage_for_failure(self) -> None:
"""
@ -2228,12 +2266,16 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
"""Assume most recent usage chunk has total usage uptil then."""
prompt_tokens: int = 0
completion_tokens: int = 0
latest_usage_chunk = None
for chunk in chunks:
if "usage" in chunk and chunk["usage"] is not None:
if "prompt_tokens" in chunk["usage"]:
prompt_tokens = chunk["usage"].get("prompt_tokens", 0) or 0
if "completion_tokens" in chunk["usage"]:
completion_tokens = chunk["usage"].get("completion_tokens", 0) or 0
usage = chunk["usage"]
latest_usage_chunk = usage
if "prompt_tokens" in usage:
prompt_tokens = usage.get("prompt_tokens", 0) or 0
if "completion_tokens" in usage:
completion_tokens = usage.get("completion_tokens", 0) or 0
returned_usage_chunk = Usage(
prompt_tokens=prompt_tokens,
@ -2241,6 +2283,15 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
total_tokens=prompt_tokens + completion_tokens,
)
if latest_usage_chunk is not None:
latest_cost = (
latest_usage_chunk.get("cost")
if isinstance(latest_usage_chunk, dict)
else getattr(latest_usage_chunk, "cost", None)
)
if latest_cost is not None:
returned_usage_chunk.cost = latest_cost
return returned_usage_chunk

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