mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal
This commit is contained in:
commit
ba4a3ba057
1759 changed files with 65610 additions and 16863 deletions
|
|
@ -1029,6 +1029,8 @@ jobs:
|
|||
- *python312_image
|
||||
working_directory: ~/project
|
||||
resource_class: large
|
||||
environment:
|
||||
REQUEST_TIMEOUT: "180"
|
||||
|
||||
steps:
|
||||
- checkout
|
||||
|
|
@ -1058,7 +1060,8 @@ jobs:
|
|||
-v -x \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 \
|
||||
-n 8"
|
||||
-n 8 \
|
||||
--reruns 1 --only-rerun Timeout"
|
||||
no_output_timeout: 15m
|
||||
|
||||
# Store test results
|
||||
|
|
|
|||
48
.github/actions/detect-backend-changes/action.yml
vendored
Normal file
48
.github/actions/detect-backend-changes/action.yml
vendored
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
name: "Detect backend-relevant changes"
|
||||
description: >-
|
||||
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
|
||||
and expose decision=run|skip. decision=skip means only ui/**, **.md or **.mdx files
|
||||
changed, so callers can short-circuit expensive steps while the job still completes
|
||||
successfully and satisfies its required status check. The decision defaults to run for
|
||||
any non pull_request event or whenever the changed set cannot be resolved, so tests are
|
||||
never skipped when the classification is uncertain.
|
||||
|
||||
outputs:
|
||||
decision:
|
||||
description: "run when backend-relevant files changed, otherwise skip"
|
||||
value: ${{ steps.classify.outputs.decision }}
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
steps:
|
||||
- id: classify
|
||||
shell: bash
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
set -uo pipefail
|
||||
if [ -z "${BASE_SHA:-}" ]; then
|
||||
echo "detect-backend-changes: not a pull_request event; running job"
|
||||
echo "decision=run" >> "${GITHUB_OUTPUT}"
|
||||
exit 0
|
||||
fi
|
||||
if ! git fetch --no-tags --depth=1 origin "${BASE_SHA}" >/dev/null 2>&1; then
|
||||
echo "detect-backend-changes: could not fetch base ${BASE_SHA}; running job"
|
||||
echo "decision=run" >> "${GITHUB_OUTPUT}"
|
||||
exit 0
|
||||
fi
|
||||
changed="$(git diff --name-only "${BASE_SHA}" HEAD 2>/dev/null)" || {
|
||||
echo "detect-backend-changes: git diff failed; running job"
|
||||
echo "decision=run" >> "${GITHUB_OUTPUT}"
|
||||
exit 0
|
||||
}
|
||||
if [ -z "${changed}" ]; then
|
||||
echo "detect-backend-changes: no changed files vs ${BASE_SHA}; skipping job"
|
||||
echo "decision=skip" >> "${GITHUB_OUTPUT}"
|
||||
exit 0
|
||||
fi
|
||||
echo "detect-backend-changes: changed files vs ${BASE_SHA}:"
|
||||
printf '%s\n' "${changed}" | sed 's/^/ /'
|
||||
decision="$(printf '%s\n' "${changed}" | bash .circleci/scripts/classify_changes.sh backend)" || decision="run"
|
||||
echo "detect-backend-changes: decision=${decision}"
|
||||
echo "decision=${decision}" >> "${GITHUB_OUTPUT}"
|
||||
24
.github/pull_request_template.md
vendored
24
.github/pull_request_template.md
vendored
|
|
@ -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
|
||||
|
|
|
|||
13
.github/workflows/_test-unit-base.yml
vendored
13
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -45,12 +45,18 @@ jobs:
|
|||
name: Run tests
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
outputs:
|
||||
decision: ${{ steps.changes.outputs.decision }}
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Detect backend-relevant changes
|
||||
id: changes
|
||||
uses: ./.github/actions/detect-backend-changes
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -72,16 +78,19 @@ jobs:
|
|||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
||||
- name: Generate Prisma client
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
env:
|
||||
TEST_PATH: ${{ inputs.test-path }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
|
|
@ -114,7 +123,7 @@ jobs:
|
|||
fi
|
||||
|
||||
- name: Save coverage report
|
||||
if: always()
|
||||
if: always() && steps.changes.outputs.decision != 'skip'
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
|
|
@ -124,7 +133,7 @@ jobs:
|
|||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always()
|
||||
if: always() && needs.run.outputs.decision != 'skip'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
61
.github/workflows/create_daily_oss_branch.yml
vendored
Normal file
61
.github/workflows/create_daily_oss_branch.yml
vendored
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
name: Create Daily OSS Branch
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: "0 16 * * 1-5" # 9am PT during daylight saving time, weekdays.
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
date:
|
||||
description: "Branch date in YYYY_MM_DD format. Defaults to today's UTC date."
|
||||
required: false
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
jobs:
|
||||
create-oss-branch:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Create dated OSS branch
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
REQUESTED_DATE: ${{ inputs.date }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
|
||||
if [ -n "${REQUESTED_DATE}" ]; then
|
||||
if ! echo "${REQUESTED_DATE}" | grep -Eq '^[0-9]{4}_[0-9]{2}_[0-9]{2}$'; then
|
||||
echo "::error::date must use YYYY_MM_DD format, got '${REQUESTED_DATE}'"
|
||||
exit 1
|
||||
fi
|
||||
BRANCH_DATE="${REQUESTED_DATE}"
|
||||
else
|
||||
BRANCH_DATE="$(date -u +'%Y_%m_%d')"
|
||||
fi
|
||||
|
||||
BRANCH_NAME="litellm_oss_daily_${BRANCH_DATE}"
|
||||
echo "Creating branch: ${BRANCH_NAME}"
|
||||
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
|
||||
git fetch origin main "${BRANCH_NAME}" || true
|
||||
|
||||
if git show-ref --verify --quiet "refs/remotes/origin/${BRANCH_NAME}"; then
|
||||
echo "Branch ${BRANCH_NAME} already exists. Skipping creation."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
git checkout -b "${BRANCH_NAME}" origin/main
|
||||
git push "https://x-access-token:${GITHUB_TOKEN}@github.com/${GITHUB_REPOSITORY}.git" "${BRANCH_NAME}"
|
||||
echo "Successfully created and pushed branch: ${BRANCH_NAME}"
|
||||
4
.github/workflows/guard-main-branch.yml
vendored
4
.github/workflows/guard-main-branch.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
50
.github/workflows/oss_daily_guardrails.yml
vendored
Normal file
50
.github/workflows/oss_daily_guardrails.yml
vendored
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
name: OSS Daily Guardrails
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- "litellm_oss_daily_20*"
|
||||
pull_request:
|
||||
branches:
|
||||
- "litellm_oss_daily_20*"
|
||||
- litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
oss-safe-checks:
|
||||
name: Run OSS daily safe checks
|
||||
if: startsWith(github.ref_name, 'litellm_oss_daily_20') || startsWith(github.head_ref, 'litellm_oss_daily_20') || startsWith(github.base_ref, 'litellm_oss_daily_20')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Run secret scan test
|
||||
run: |
|
||||
uv run --frozen --with 'pytest==9.0.2' pytest tests/litellm/test_no_hardcoded_secrets.py -v
|
||||
|
||||
- name: Run Ruff
|
||||
run: |
|
||||
uv sync --frozen
|
||||
cd litellm
|
||||
uv run --no-sync ruff check .
|
||||
12
.github/workflows/test-linting.yml
vendored
12
.github/workflows/test-linting.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
76
.github/workflows/test-litellm-ui-build.yml
vendored
76
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
92
.github/workflows/test-litellm-ui-lint.yml
vendored
Normal file
92
.github/workflows/test-litellm-ui-lint.yml
vendored
Normal 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
|
||||
|
|
@ -32,6 +32,10 @@ jobs:
|
|||
path: docs/my-website
|
||||
persist-credentials: false
|
||||
|
||||
- name: Detect backend-relevant changes
|
||||
id: changes
|
||||
uses: ./.github/actions/detect-backend-changes
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -53,10 +57,12 @@ jobs:
|
|||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
||||
- name: Generate Prisma client
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
|
|
@ -64,6 +70,7 @@ jobs:
|
|||
|
||||
# Run the same documentation tests that CircleCI ran (as direct Python scripts)
|
||||
- name: Run documentation validation tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
|
||||
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
|
||||
|
|
|
|||
7
.github/workflows/test-unit-proxy-legacy.yml
vendored
7
.github/workflows/test-unit-proxy-legacy.yml
vendored
|
|
@ -49,6 +49,10 @@ jobs:
|
|||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Detect backend-relevant changes
|
||||
id: changes
|
||||
uses: ./.github/actions/detect-backend-changes
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -70,16 +74,19 @@ jobs:
|
|||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
||||
- name: Generate Prisma client
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests - ${{ matrix.test-group.name }}
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
env:
|
||||
TEST_PATH: ${{ matrix.test-group.path }}
|
||||
run: |
|
||||
|
|
|
|||
23
.github/workflows/test_server_root_path.yml
vendored
23
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -16,6 +16,7 @@ jobs:
|
|||
timeout-minutes: 30
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
root_path: ["/api/v1", "/llmproxy"]
|
||||
|
||||
|
|
@ -108,8 +109,26 @@ jobs:
|
|||
- name: Install UI deps and Chromium
|
||||
working-directory: ui/litellm-dashboard
|
||||
run: |
|
||||
npm ci
|
||||
npx playwright install --with-deps chromium
|
||||
retry() {
|
||||
local attempt=1
|
||||
local max_attempts=4
|
||||
until "$@"; do
|
||||
if [ "$attempt" -ge "$max_attempts" ]; then
|
||||
echo "Command failed after $attempt attempts: $*"
|
||||
return 1
|
||||
fi
|
||||
echo "Attempt $attempt failed: $*. Retrying in $((attempt * 15))s..."
|
||||
sleep $((attempt * 15))
|
||||
attempt=$((attempt + 1))
|
||||
done
|
||||
}
|
||||
|
||||
npm config set fetch-retries 5
|
||||
npm config set fetch-retry-mintimeout 20000
|
||||
npm config set fetch-retry-maxtimeout 120000
|
||||
|
||||
retry npm ci
|
||||
retry npx playwright install --with-deps chromium
|
||||
|
||||
- name: Run SERVER_ROOT_PATH redirect e2e
|
||||
working-directory: ui/litellm-dashboard
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
28
Makefile
28
Makefile
|
|
@ -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
|
||||
|
|
|
|||
13
README.md
13
README.md
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/fallback",
|
||||
"/fallbacks",
|
||||
"/cache_settings",
|
||||
"/coordination_redis/",
|
||||
"/cost_tracking",
|
||||
"/cost/",
|
||||
"/credentials",
|
||||
|
|
|
|||
|
|
@ -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": 32151
|
||||
"limit": 32141
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 177
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 7
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 1212
|
||||
"limit": 1209
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 165
|
||||
|
|
|
|||
|
|
@ -6,4 +6,4 @@ Code in this folder is licensed under a commercial license. Please review the [L
|
|||
|
||||
👉 **Using in an Enterprise / Need specific features ?** Meet with us [here](https://enterprise.litellm.ai/demo?month=2024-02)
|
||||
|
||||
See all Enterprise Features here 👉 [Docs](https://docs.litellm.ai/docs/proxy/enterprise)
|
||||
See all Enterprise Features here 👉 [Docs](https://docs.litellm.ai/docs/enterprise)
|
||||
|
|
|
|||
|
|
@ -919,9 +919,9 @@ class BaseEmailLogger(CustomLogger):
|
|||
"""
|
||||
Construct invitation link for the user
|
||||
|
||||
# http://localhost:4000/ui?invitation_id=7a096b3a-37c6-440f-9dd1-ba22e8043f6b
|
||||
# http://localhost:4000/ui/onboarding?invitation_id=7a096b3a-37c6-440f-9dd1-ba22e8043f6b
|
||||
"""
|
||||
return f"{base_url}/ui?invitation_id={invitation_id}"
|
||||
return f"{base_url}/ui/onboarding?invitation_id={invitation_id}"
|
||||
|
||||
async def send_email(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.48"
|
||||
version = "0.1.49"
|
||||
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.48"
|
||||
version = "0.1.49"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -76,10 +76,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 "-") -}}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
143
helm/litellm-helm/tests/coordination_redis_tests.yaml
Normal file
143
helm/litellm-helm/tests/coordination_redis_tests.yaml
Normal 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"
|
||||
|
|
@ -331,12 +331,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:
|
||||
|
|
|
|||
|
|
@ -213,6 +213,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 +230,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 }}
|
||||
|
|
|
|||
109
helm/litellm/tests/redis_env_tests.yaml
Normal file
109
helm/litellm/tests/redis_env_tests.yaml
Normal 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
|
||||
|
|
@ -100,7 +100,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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,8 @@
|
|||
-- Timestamp sorts before some already-applied migrations; this is safe: the
|
||||
-- runner is `prisma migrate deploy`, which applies every pending migration
|
||||
-- regardless of name order (utils.py has an informational check for exactly
|
||||
-- this), and IF NOT EXISTS keeps a re-apply idempotent.
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "token_exchange_endpoint" TEXT;
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "audience" TEXT;
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "subject_token_type" TEXT;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "token_exchange_profile" TEXT;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "dcr_bridge" BOOLEAN;
|
||||
|
|
@ -329,10 +329,17 @@ model LiteLLM_MCPServerTable {
|
|||
token_url String?
|
||||
registration_url String?
|
||||
oauth2_flow String?
|
||||
token_exchange_endpoint String?
|
||||
// Named for the RFC 8693 "audience" token-exchange request parameter (that flow only).
|
||||
// RFC 8707 resource indicators are a separate concept, named "resource" in the v2 egress types.
|
||||
audience String?
|
||||
subject_token_type String?
|
||||
token_exchange_profile String?
|
||||
allow_all_keys Boolean @default(false)
|
||||
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?
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.75"
|
||||
version = "0.4.76"
|
||||
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.76"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
@ -678,9 +689,8 @@ def get_redis_connection_pool(
|
|||
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
|
||||
|
||||
connection_class = async_redis.Connection
|
||||
if "ssl" in redis_kwargs:
|
||||
if redis_kwargs.pop("ssl", False):
|
||||
connection_class = async_redis.SSLConnection
|
||||
redis_kwargs.pop("ssl", None)
|
||||
redis_kwargs["connection_class"] = connection_class
|
||||
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@
|
|||
"computer-use-2025-01-24": "computer-use-2025-01-24",
|
||||
"computer-use-2025-11-24": "computer-use-2025-11-24",
|
||||
"context-1m-2025-08-07": "context-1m-2025-08-07",
|
||||
"context-management-2025-06-27": null,
|
||||
"context-management-2025-06-27": "context-management-2025-06-27",
|
||||
"effort-2025-11-24": "effort-2025-11-24",
|
||||
"fast-mode-2026-02-01": null,
|
||||
"files-api-2025-04-14": null,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -715,6 +715,7 @@ openai_compatible_endpoints: List = [
|
|||
"https://api.clarifai.com/v2/ext/openai/v1",
|
||||
"https://api.libertai.io/v1",
|
||||
"https://pinstripes.io/v1",
|
||||
"https://api.meta.ai/v1",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -781,6 +782,7 @@ openai_compatible_providers: List = [
|
|||
"ragflow",
|
||||
"pinstripes", # Pinstripes - JSON-configured provider
|
||||
"darkbloom",
|
||||
"meta", # Meta Model API (Muse Spark) - JSON-configured provider
|
||||
]
|
||||
openai_text_completion_compatible_providers: List = [ # providers that support `/v1/completions`
|
||||
"together_ai",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1180,20 +1180,6 @@ class ModifyResponseException(Exception):
|
|||
super().__init__(message)
|
||||
|
||||
|
||||
class GuardrailInterventionNormalStringError(
|
||||
Exception
|
||||
): # custom exception to raise when a guardrail intervenes, but we want to return a normal string to the user
|
||||
def __init__(self, message: str):
|
||||
self.message = message
|
||||
super().__init__(self.message)
|
||||
|
||||
def __str__(self):
|
||||
return self.message
|
||||
|
||||
def __repr__(self):
|
||||
return self.__str__()
|
||||
|
||||
|
||||
class SensitiveDataRouteException(Exception):
|
||||
"""
|
||||
Exception raised when a guardrail detects sensitive data and wants to reroute the request.
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]],
|
||||
|
|
@ -757,6 +795,12 @@ class CustomGuardrail(CustomLogger):
|
|||
# raw provider JSON so redaction is not duplicated upstream).
|
||||
clean_guardrail_response = redact_nested_match_and_regex_keys(clean_guardrail_response)
|
||||
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import (
|
||||
mask_credentials_in_payload,
|
||||
)
|
||||
|
||||
clean_guardrail_response = mask_credentials_in_payload(clean_guardrail_response)
|
||||
|
||||
slg = StandardLoggingGuardrailInformation(
|
||||
guardrail_name=self.guardrail_name,
|
||||
guardrail_provider=guardrail_provider,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ from litellm.integrations.otel.model.semconv import (
|
|||
GenAIProvider,
|
||||
JsonRpc,
|
||||
LiteLLM,
|
||||
LiteLLMError,
|
||||
MCPMethod,
|
||||
Metric,
|
||||
Network,
|
||||
|
|
@ -87,6 +88,7 @@ __all__ = [
|
|||
"HTTP",
|
||||
"JsonRpc",
|
||||
"LiteLLM",
|
||||
"LiteLLMError",
|
||||
"MCP",
|
||||
"MCPMethod",
|
||||
"Metric",
|
||||
|
|
|
|||
|
|
@ -16,9 +16,11 @@ from litellm.integrations.otel.model.payloads import (
|
|||
MCPListToolsSpanData,
|
||||
MCPToolCallSpanData,
|
||||
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
|
||||
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError
|
||||
from litellm.integrations.otel.model.spans import (
|
||||
SPAN_REGISTRY,
|
||||
SpanRole,
|
||||
|
|
@ -49,15 +51,38 @@ _NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = {
|
|||
_DEDUP_CACHE_MAX = 10_000
|
||||
|
||||
|
||||
def _stamp_otel_error_attributes(span: Span, error_type: str, resolved_message: str) -> None:
|
||||
"""Stamp the OTel-semconv error attributes (``error.type`` + ``error.message``).
|
||||
``error_type`` and ``resolved_message`` are ``finish_span``'s already-computed
|
||||
fallback chains, so the pair on the status, event, and attributes stays in
|
||||
lockstep."""
|
||||
span.set_attribute(Error.TYPE, error_type)
|
||||
span.set_attribute(Error.MESSAGE, resolved_message)
|
||||
|
||||
|
||||
def _stamp_litellm_error_attributes(span: Span, error: SpanError) -> None:
|
||||
"""Stamp litellm-specific error detail attributes. Emitted only when the
|
||||
corresponding field is populated so guardrail-shape errors carrying only a
|
||||
message aren't polluted with empty detail keys."""
|
||||
if error.code:
|
||||
span.set_attribute(LiteLLMError.CODE, error.code)
|
||||
if error.stack_trace:
|
||||
span.set_attribute(LiteLLMError.STACK_TRACE, error.stack_trace)
|
||||
if error.llm_provider:
|
||||
span.set_attribute(LiteLLMError.LLM_PROVIDER, error.llm_provider)
|
||||
|
||||
|
||||
class SpanEmitter:
|
||||
def __init__(
|
||||
self,
|
||||
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] = (
|
||||
|
|
@ -190,16 +215,25 @@ class SpanEmitter:
|
|||
if error and (error.error_type or error.message):
|
||||
error_type = error.error_type or "error"
|
||||
message = error.message or error.error_type or "error"
|
||||
span.set_attribute(Error.TYPE, error_type)
|
||||
_stamp_otel_error_attributes(span, error_type, message)
|
||||
_stamp_litellm_error_attributes(span, error)
|
||||
span.set_status(Status(StatusCode.ERROR, message))
|
||||
# Carry the full message on the standard ``exception`` event so backends
|
||||
# map it as full text under ``exception.message``. Setting it as a bare
|
||||
# string attribute instead lets backends like Elasticsearch dynamic-map
|
||||
# it to a ``keyword`` capped at 1024 chars, truncating the message.
|
||||
# Also emit the semconv ``exception`` event so backends that
|
||||
# dynamic-map unknown string span attrs to ``keyword`` (e.g.
|
||||
# Elasticsearch with a 1024-char ``ignore_above``) still see the
|
||||
# full untruncated message on the recognized event field.
|
||||
span.add_event(
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ====================================================================== #
|
||||
|
|
|
|||
|
|
@ -141,6 +141,9 @@ class LLMCost:
|
|||
class SpanError:
|
||||
error_type: str | None = None
|
||||
message: str | None = None
|
||||
code: str | None = None
|
||||
stack_trace: str | None = None
|
||||
llm_provider: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -571,6 +574,9 @@ def _parse_error(payload: "StandardLoggingPayload") -> SpanError | None:
|
|||
return SpanError(
|
||||
error_type=as_str(info.get("error_class")) or as_str(info.get("error_code")),
|
||||
message=as_str(info.get("error_message")) or as_str(payload.get("error_str")),
|
||||
code=as_str(info.get("error_code")),
|
||||
stack_trace=as_str(info.get("traceback")),
|
||||
llm_provider=as_str(info.get("llm_provider")),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -144,7 +144,24 @@ class Client:
|
|||
|
||||
|
||||
class Error:
|
||||
"""OTel-defined error attribute keys, from the semconv ``error.*`` registry.
|
||||
``MESSAGE`` is marked *Deprecated* upstream in favor of domain-specific
|
||||
error message keys plus ``exception.message`` on the exception event, but
|
||||
litellm still stamps it."""
|
||||
|
||||
TYPE: Final = "error.type"
|
||||
MESSAGE: Final = "error.message"
|
||||
|
||||
|
||||
class LiteLLMError:
|
||||
"""Detail keys for the mapped provider exception of a failed LLM call.
|
||||
OTel semconv does not define these, so they live under the ``litellm.*``
|
||||
vendor namespace rather than squatting on the semconv-owned ``error.*``
|
||||
namespace."""
|
||||
|
||||
CODE: Final = "litellm.provider.error.code"
|
||||
STACK_TRACE: Final = "litellm.provider.error.stack_trace"
|
||||
LLM_PROVIDER: Final = "litellm.provider.error.llm_provider"
|
||||
|
||||
|
||||
class ExceptionEvent:
|
||||
|
|
@ -160,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:
|
||||
|
|
|
|||
52
litellm/integrations/otel/plumbing/events.py
Normal file
52
litellm/integrations/otel/plumbing/events.py
Normal 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,
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1618,6 +1618,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)
|
||||
|
|
@ -3332,6 +3340,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 +3464,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 +3596,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,
|
||||
|
|
@ -3619,6 +3636,7 @@ class PrometheusLogger(CustomLogger):
|
|||
hashed_token=user_api_key_dict.token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_cache_only=True,
|
||||
)
|
||||
if key_object:
|
||||
user_api_key_dict.budget_reset_at = key_object.budget_reset_at
|
||||
|
|
@ -3641,6 +3659,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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -346,6 +346,9 @@ def get_llm_provider(
|
|||
elif endpoint == "https://pinstripes.io/v1":
|
||||
custom_llm_provider = "pinstripes"
|
||||
dynamic_api_key = get_secret_str("PINSTRIPES_API_KEY")
|
||||
elif endpoint == "https://api.meta.ai/v1":
|
||||
custom_llm_provider = "meta"
|
||||
dynamic_api_key = get_secret_str("META_API_KEY")
|
||||
|
||||
if api_base is not None and not isinstance(api_base, str):
|
||||
raise Exception("api base needs to be a string. api_base={}".format(api_base))
|
||||
|
|
@ -471,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:
|
||||
|
|
|
|||
|
|
@ -95,6 +95,17 @@ class HealthCheckHelpers:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
logging_obj = filtered_model_params.get("litellm_logging_obj")
|
||||
if logging_obj is not None:
|
||||
api_base = filtered_model_params.get("api_base")
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=filtered_model_params,
|
||||
model=filtered_model_params.get("model"),
|
||||
user=None,
|
||||
optional_params={},
|
||||
litellm_params={"api_base": api_base} if api_base else None,
|
||||
)
|
||||
|
||||
if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.alist_batches(**filtered_model_params)
|
||||
else:
|
||||
|
|
@ -188,6 +199,7 @@ class HealthCheckHelpers:
|
|||
api_base=model_params.get("api_base", None),
|
||||
api_key=model_params.get("api_key", None),
|
||||
api_version=model_params.get("api_version", None),
|
||||
model_params=model_params,
|
||||
),
|
||||
"batch": lambda: HealthCheckHelpers._batch_health_check(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
139
litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py
Normal file
139
litellm/litellm_core_utils/llm_cost_calc/tiered_pricing.py
Normal 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)
|
||||
|
|
@ -3626,6 +3626,7 @@ class BedrockImageProcessor:
|
|||
|
||||
def _convert_to_bedrock_tool_call_invoke(
|
||||
tool_calls: list,
|
||||
model: Optional[str] = None,
|
||||
) -> List[BedrockContentBlock]:
|
||||
"""
|
||||
OpenAI tool invokes:
|
||||
|
|
@ -3701,7 +3702,13 @@ def _convert_to_bedrock_tool_call_invoke(
|
|||
# cache_control applies to the whole original
|
||||
# tool call; attach after the last split block.
|
||||
if tool.get("cache_control", None) is not None:
|
||||
_parts_list.append(BedrockContentBlock(cachePoint=CachePointBlock(type="default")))
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
{"cache_control": tool["cache_control"]},
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
_parts_list.append(_cache_point_block)
|
||||
continue
|
||||
# Fallback: no objects extracted — use empty dict.
|
||||
arguments_dict = {}
|
||||
|
|
@ -3712,8 +3719,13 @@ def _convert_to_bedrock_tool_call_invoke(
|
|||
|
||||
# Check for cache_control and add a separate cachePoint block
|
||||
if tool.get("cache_control", None) is not None:
|
||||
cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default"))
|
||||
_parts_list.append(cache_point_block)
|
||||
cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
{"cache_control": tool["cache_control"]},
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if cache_point_block is not None:
|
||||
_parts_list.append(cache_point_block)
|
||||
return _parts_list
|
||||
except Exception as e:
|
||||
raise Exception(
|
||||
|
|
@ -4377,6 +4389,7 @@ class BedrockConverseMessagesProcessor:
|
|||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block=cast(OpenAIMessageContentListBlock, element),
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
_parts.append(_cache_point_block)
|
||||
|
|
@ -4384,7 +4397,7 @@ class BedrockConverseMessagesProcessor:
|
|||
elif message_block["content"] and isinstance(message_block["content"], str):
|
||||
_part = BedrockContentBlock(text=messages[msg_i]["content"])
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block, block_type="content_block"
|
||||
message_block, block_type="content_block", model=model
|
||||
)
|
||||
user_content.append(_part)
|
||||
if _cache_point_block is not None:
|
||||
|
|
@ -4416,22 +4429,27 @@ class BedrockConverseMessagesProcessor:
|
|||
tool_content.append(tool_call_result)
|
||||
|
||||
# Check if we need to add a separate cachePoint block
|
||||
has_cache_control = False
|
||||
tool_msg_cache_control = None
|
||||
|
||||
# Check for message-level cache_control
|
||||
if current_message.get("cache_control", None) is not None:
|
||||
has_cache_control = True
|
||||
tool_msg_cache_control = current_message["cache_control"]
|
||||
# Check for content-level cache_control in list content
|
||||
elif isinstance(current_message.get("content"), list):
|
||||
for content_element in current_message["content"]:
|
||||
if isinstance(content_element, dict) and content_element.get("cache_control", None) is not None:
|
||||
has_cache_control = True
|
||||
tool_msg_cache_control = content_element["cache_control"]
|
||||
break
|
||||
|
||||
# Add a separate cachePoint block if cache_control is present
|
||||
if has_cache_control:
|
||||
cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default"))
|
||||
tool_content.append(cache_point_block)
|
||||
if tool_msg_cache_control is not None:
|
||||
cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
{"cache_control": tool_msg_cache_control},
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if cache_point_block is not None:
|
||||
tool_content.append(cache_point_block)
|
||||
|
||||
msg_i += 1
|
||||
# Deduplicate toolResult blocks with the same toolUseId
|
||||
|
|
@ -4509,6 +4527,7 @@ class BedrockConverseMessagesProcessor:
|
|||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block=cast(OpenAIMessageContentListBlock, element),
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
assistants_parts.append(_cache_point_block)
|
||||
|
|
@ -4520,14 +4539,14 @@ class BedrockConverseMessagesProcessor:
|
|||
# If content is empty/whitespace, skip it (don't add a placeholder)
|
||||
# Add cache point block for assistant string content
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
assistant_message_block, block_type="content_block"
|
||||
assistant_message_block, block_type="content_block", model=model
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
assistant_content.append(_cache_point_block)
|
||||
|
||||
_tool_calls = assistant_message_block.get("tool_calls", [])
|
||||
if _tool_calls:
|
||||
assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls))
|
||||
assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls, model=model))
|
||||
|
||||
msg_i += 1
|
||||
|
||||
|
|
@ -4745,6 +4764,7 @@ def _bedrock_converse_messages_pt(
|
|||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block=cast(OpenAIMessageContentListBlock, element),
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
_parts.append(_cache_point_block)
|
||||
|
|
@ -4752,7 +4772,7 @@ def _bedrock_converse_messages_pt(
|
|||
elif message_block["content"] and isinstance(message_block["content"], str):
|
||||
_part = BedrockContentBlock(text=messages[msg_i]["content"])
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block, block_type="content_block"
|
||||
message_block, block_type="content_block", model=model
|
||||
)
|
||||
user_content.append(_part)
|
||||
if _cache_point_block is not None:
|
||||
|
|
@ -4786,22 +4806,27 @@ def _bedrock_converse_messages_pt(
|
|||
tool_content.append(tool_call_result)
|
||||
|
||||
# Check if we need to add a separate cachePoint block
|
||||
has_cache_control = False
|
||||
tool_msg_cache_control = None
|
||||
|
||||
# Check for message-level cache_control
|
||||
if current_message.get("cache_control", None) is not None:
|
||||
has_cache_control = True
|
||||
tool_msg_cache_control = current_message["cache_control"]
|
||||
# Check for content-level cache_control in list content
|
||||
elif isinstance(current_message.get("content"), list):
|
||||
for content_element in current_message["content"]:
|
||||
if isinstance(content_element, dict) and content_element.get("cache_control", None) is not None:
|
||||
has_cache_control = True
|
||||
tool_msg_cache_control = content_element["cache_control"]
|
||||
break
|
||||
|
||||
# Add a separate cachePoint block if cache_control is present
|
||||
if has_cache_control:
|
||||
cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default"))
|
||||
tool_content.append(cache_point_block)
|
||||
if tool_msg_cache_control is not None:
|
||||
cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
{"cache_control": tool_msg_cache_control},
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if cache_point_block is not None:
|
||||
tool_content.append(cache_point_block)
|
||||
|
||||
msg_i += 1
|
||||
# Deduplicate toolResult blocks with the same toolUseId
|
||||
|
|
@ -4882,6 +4907,7 @@ def _bedrock_converse_messages_pt(
|
|||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
message_block=cast(OpenAIMessageContentListBlock, element),
|
||||
block_type="content_block",
|
||||
model=model,
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
assistants_parts.append(_cache_point_block)
|
||||
|
|
@ -4892,13 +4918,13 @@ def _bedrock_converse_messages_pt(
|
|||
assistant_content.append(BedrockContentBlock(text=_assistant_content))
|
||||
# Add cache point block for assistant string content
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
assistant_message_block, block_type="content_block"
|
||||
assistant_message_block, block_type="content_block", model=model
|
||||
)
|
||||
if _cache_point_block is not None:
|
||||
assistant_content.append(_cache_point_block)
|
||||
_tool_calls = assistant_message_block.get("tool_calls", [])
|
||||
if _tool_calls:
|
||||
assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls))
|
||||
assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls, model=model))
|
||||
|
||||
msg_i += 1
|
||||
|
||||
|
|
@ -5468,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
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Dict, List, Optional, Set
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
|
||||
|
||||
|
||||
|
|
@ -153,6 +155,39 @@ def mask_sensitive_structure(data: object) -> object:
|
|||
return _error_masker.mask(data)
|
||||
|
||||
|
||||
def mask_credentials_in_payload(data: object) -> object:
|
||||
"""Return a copy of ``data`` where string values under sensitive-named keys
|
||||
are masked but every other value (``None``, ``int``, ``float``, ``bool``,
|
||||
``bytes``, ``datetime``, tuples, sets, typed objects) is preserved by
|
||||
identity, and dicts/lists are rebuilt structurally.
|
||||
|
||||
Use this for logging payloads that carry response data through to
|
||||
SpendLogs / OTel / Langfuse, where :meth:`SensitiveDataMasker.mask`'s
|
||||
config-dump semantics (``None`` -> ``"None"``, tuples stringified,
|
||||
objects flattened via ``__dict__``) would silently distort the record.
|
||||
|
||||
Sensitive-key detection is delegated to the shared
|
||||
:class:`SensitiveDataMasker` so pattern updates stay in one place.
|
||||
"""
|
||||
return _walk_payload(data, key_is_sensitive=False, depth=0)
|
||||
|
||||
|
||||
def _walk_payload(node: object, key_is_sensitive: bool, depth: int) -> object:
|
||||
if depth >= DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER:
|
||||
return node
|
||||
if isinstance(node, Mapping):
|
||||
return {k: _walk_payload(v, _default_masker.is_sensitive_key(k), depth + 1) for k, v in node.items()}
|
||||
if isinstance(node, list):
|
||||
return [_walk_payload(item, key_is_sensitive, depth + 1) for item in node]
|
||||
if isinstance(node, tuple):
|
||||
return tuple(_walk_payload(item, key_is_sensitive, depth + 1) for item in node)
|
||||
if isinstance(node, BaseModel):
|
||||
return _walk_payload(node.model_dump(), key_is_sensitive, depth)
|
||||
if key_is_sensitive and isinstance(node, str) and node:
|
||||
return _default_masker._mask_value(node)
|
||||
return node
|
||||
|
||||
|
||||
def mask_sensitive_keys(data: Dict[str, Any], sensitive_fields: Set[str]) -> Dict[str, Any]:
|
||||
"""Return a new dict with values masked for keys listed in ``sensitive_fields``.
|
||||
|
||||
|
|
|
|||
|
|
@ -227,6 +227,10 @@ DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING = (
|
|||
"Sonnet 4.6+, and Mythos Preview."
|
||||
)
|
||||
|
||||
DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING = (
|
||||
"Dropping adaptive `thinking` for model=%s: max_tokens is too small to fit the minimum thinking budget."
|
||||
)
|
||||
|
||||
DROP_UNSUPPORTED_SPEED_WARNING = (
|
||||
"Dropping unsupported `speed` for model=%s (drop_params=True). Fast mode is only supported on select Opus models."
|
||||
)
|
||||
|
|
@ -266,6 +270,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "anthropic"
|
||||
|
||||
@property
|
||||
def _resolved_provider(self) -> str:
|
||||
return self.custom_llm_provider or "anthropic"
|
||||
|
||||
@classmethod
|
||||
def get_config(cls, *, model: Optional[str] = None):
|
||||
config = super().get_config()
|
||||
|
|
@ -335,23 +343,26 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return any(v in model_lower for v in ("opus-4-7", "opus_4_7", "opus-4.7", "opus_4.7"))
|
||||
|
||||
@staticmethod
|
||||
def _supports_effort_level(model: str, level: str) -> bool:
|
||||
def _supports_effort_level(model: str, level: str, custom_llm_provider: str) -> bool:
|
||||
"""Check ``supports_{level}_reasoning_effort`` in the model map."""
|
||||
return AnthropicConfig._supports_model_capability(model, f"supports_{level}_reasoning_effort")
|
||||
return AnthropicConfig._supports_model_capability(
|
||||
model, f"supports_{level}_reasoning_effort", custom_llm_provider
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
|
||||
def _validate_effort_for_model(model: str, effort: Optional[str], custom_llm_provider: str) -> Optional[str]:
|
||||
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
|
||||
if effort == "max" and not (
|
||||
AnthropicConfig._is_adaptive_thinking_model(model) or AnthropicConfig._supports_effort_level(model, "max")
|
||||
AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider)
|
||||
or AnthropicConfig._supports_effort_level(model, "max", custom_llm_provider)
|
||||
):
|
||||
return f"effort='max' is not supported by this model. Got model: {model}"
|
||||
if effort == "xhigh" and not AnthropicConfig._supports_effort_level(model, "xhigh"):
|
||||
if effort == "xhigh" and not AnthropicConfig._supports_effort_level(model, "xhigh", custom_llm_provider):
|
||||
return f"effort='xhigh' is not supported by this model. Got model: {model}"
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _model_supports_effort_param(model: str) -> bool:
|
||||
def _model_supports_effort_param(model: str, custom_llm_provider: str) -> bool:
|
||||
"""Whether the model accepts ``output_config.effort`` at all.
|
||||
|
||||
A model qualifies if its map entry advertises ``supports_output_config``
|
||||
|
|
@ -359,10 +370,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
signals: e.g. Claude Opus 4.5 supports ``output_config`` without
|
||||
advertising a non-default (max/xhigh) effort level.
|
||||
"""
|
||||
if AnthropicConfig._supports_model_capability(model, "supports_output_config"):
|
||||
if AnthropicConfig._supports_model_capability(model, "supports_output_config", custom_llm_provider):
|
||||
return True
|
||||
return any(
|
||||
AnthropicConfig._supports_effort_level(model, level)
|
||||
AnthropicConfig._supports_effort_level(model, level, custom_llm_provider)
|
||||
for level in ("low", "minimal", "medium", "high", "xhigh", "max")
|
||||
)
|
||||
|
||||
|
|
@ -451,7 +462,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
if (
|
||||
"claude-3-7-sonnet" in model
|
||||
or AnthropicConfig._is_adaptive_thinking_model(model)
|
||||
or AnthropicConfig._is_adaptive_thinking_model(model, self._resolved_provider)
|
||||
or supports_reasoning(
|
||||
model=model,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
|
|
@ -1159,11 +1170,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _map_reasoning_effort(
|
||||
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
llm_provider: str = "anthropic",
|
||||
) -> Optional[AnthropicThinkingParam]:
|
||||
"""Capability probes read the cost map under ``custom_llm_provider``; ``llm_provider`` only tags raised exceptions."""
|
||||
if reasoning_effort is None or reasoning_effort == "none":
|
||||
return None
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
return AnthropicThinkingParam(
|
||||
type="adaptive",
|
||||
)
|
||||
|
|
@ -1211,6 +1224,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _cap_thinking_budget_to_max_tokens(
|
||||
thinking: AnthropicThinkingParam, max_tokens: Optional[int]
|
||||
) -> Optional[AnthropicThinkingParam]:
|
||||
"""Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic
|
||||
requires ``max_tokens > budget_tokens``). Returns the (possibly capped)
|
||||
thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the
|
||||
minimum thinking budget and thinking should be dropped."""
|
||||
budget = thinking.get("budget_tokens")
|
||||
if max_tokens is None or not isinstance(budget, int):
|
||||
return thinking
|
||||
if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS:
|
||||
return None
|
||||
if budget < max_tokens:
|
||||
return thinking
|
||||
return AnthropicThinkingParam(type=thinking.get("type", "enabled"), budget_tokens=max_tokens - 1)
|
||||
|
||||
def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]:
|
||||
if value is None:
|
||||
return None
|
||||
|
|
@ -1454,7 +1484,38 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
):
|
||||
optional_params["metadata"] = {"user_id": value}
|
||||
elif param == "thinking":
|
||||
optional_params["thinking"] = value
|
||||
if (
|
||||
isinstance(value, dict)
|
||||
and value.get("type") == "adaptive"
|
||||
and not AnthropicConfig._is_adaptive_thinking_model(model, self._resolved_provider)
|
||||
):
|
||||
# Callers (e.g. Claude Code) send adaptive thinking
|
||||
# unconditionally; translate it down to the legacy
|
||||
# `thinking={type: enabled, budget_tokens}` interface a
|
||||
# pre-4.6 model actually supports instead of forwarding a
|
||||
# shape the model will reject.
|
||||
max_tokens = non_default_params.get("max_completion_tokens") or non_default_params.get("max_tokens")
|
||||
legacy_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort="medium",
|
||||
model=model,
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
capped_thinking = (
|
||||
AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens)
|
||||
if legacy_thinking is not None
|
||||
else None
|
||||
)
|
||||
if capped_thinking is not None:
|
||||
optional_params["thinking"] = capped_thinking
|
||||
else:
|
||||
litellm.verbose_logger.warning(
|
||||
DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING,
|
||||
model,
|
||||
)
|
||||
optional_params.pop("thinking", None)
|
||||
else:
|
||||
optional_params["thinking"] = value
|
||||
elif param == "reasoning_effort":
|
||||
# Accept both string ("low") and dict ({"effort": "low",
|
||||
# "summary": "concise"}). The Responses->Chat parser keeps the
|
||||
|
|
@ -1471,20 +1532,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
mapped_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=effort_value,
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
if mapped_thinking is None:
|
||||
optional_params.pop("thinking", None)
|
||||
optional_params.pop("output_config", None)
|
||||
else:
|
||||
optional_params["thinking"] = mapped_thinking
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model, self._resolved_provider):
|
||||
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(effort_value)
|
||||
if mapped_effort is None:
|
||||
AnthropicConfig._raise_invalid_reasoning_effort(
|
||||
model=model,
|
||||
value=effort_value,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
optional_params["output_config"] = {"effort": mapped_effort}
|
||||
elif param == "web_search_options" and isinstance(value, dict):
|
||||
|
|
@ -1813,7 +1875,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
anthropic_messages = anthropic_messages_pt(
|
||||
model=model,
|
||||
messages=messages,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
raise AnthropicError(
|
||||
|
|
@ -1902,7 +1964,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
output_config = optional_params.get("output_config")
|
||||
if not output_config or not isinstance(output_config, dict):
|
||||
return
|
||||
if litellm.drop_params is True and not self._model_supports_effort_param(model):
|
||||
if litellm.drop_params is True and not self._model_supports_effort_param(model, self._resolved_provider):
|
||||
litellm.verbose_logger.warning(
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
model,
|
||||
|
|
@ -1916,14 +1978,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
raise litellm.exceptions.BadRequestError(
|
||||
message=(f"Invalid effort value: {effort!r}. Must be one of: 'high', 'medium', 'low', 'xhigh', 'max'"),
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
gate_error = self._validate_effort_for_model(model, effort)
|
||||
gate_error = self._validate_effort_for_model(model, effort, self._resolved_provider)
|
||||
if gate_error is not None:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=gate_error,
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
data["output_config"] = output_config
|
||||
|
||||
|
|
|
|||
|
|
@ -289,6 +289,13 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
status_code=400,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _strip_version_suffix(model: str) -> str:
|
||||
at = model.rfind("@")
|
||||
if at > 0:
|
||||
return model[:at]
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
def _model_map_lookup_candidates(model: str) -> List[str]:
|
||||
"""Model-map keys to try for ``model``: the id itself, the same id with a
|
||||
|
|
@ -324,6 +331,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
_DATED_RELEASE_SUFFIX_RE.sub("", cand),
|
||||
_DOTTED_VERSION_RE.sub(r"\1-\2", cand),
|
||||
_strip_bedrock_id_suffixes(cand),
|
||||
AnthropicModelInfo._strip_version_suffix(cand),
|
||||
)
|
||||
)
|
||||
return list(dict.fromkeys((*primary, *normalized)))
|
||||
|
|
@ -352,18 +360,43 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
return value if isinstance(value, bool) else None
|
||||
|
||||
@staticmethod
|
||||
def _supports_model_capability(model: str, key: str) -> bool:
|
||||
"""Check a boolean capability ``key`` in the model map.
|
||||
def _get_provider_resolved_capability(model: str, key: str, custom_llm_provider: str) -> Optional[bool]:
|
||||
"""Resolve boolean capability ``key`` for ``model`` under the caller's provider.
|
||||
|
||||
Strips bedrock/vertex prefixes so a provider-routed Claude still
|
||||
resolves to the Anthropic model-map entry.
|
||||
Returns the flag when the provider-aware lookup resolves ``model`` to an
|
||||
entry (or fallback rule) that sets it explicitly, and ``None`` when the
|
||||
model does not resolve under that provider or the resolved entry has no
|
||||
opinion on ``key``.
|
||||
"""
|
||||
from litellm.utils import _get_model_info_helper
|
||||
|
||||
try:
|
||||
resolved_model, resolved_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
value = _get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider).get(key)
|
||||
except Exception: # noqa: BLE001 # _get_model_info_helper raises bare Exception for unmapped models
|
||||
return None
|
||||
return value if isinstance(value, bool) else None
|
||||
|
||||
@staticmethod
|
||||
def _supports_model_capability(model: str, key: str, custom_llm_provider: str) -> bool:
|
||||
"""Check a boolean capability ``key`` in the model map under the caller's provider.
|
||||
|
||||
The provider-aware lookup is authoritative when it resolves an explicit flag,
|
||||
so ``key: false`` on the provider-namespaced entry wins over every fallback.
|
||||
Otherwise ``_supports_factory``'s provider-level fallbacks and the raw
|
||||
model-map walk remain as backstops for alias forms the lookup misses.
|
||||
"""
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
resolved = AnthropicModelInfo._get_provider_resolved_capability(model, key, custom_llm_provider)
|
||||
if resolved is not None:
|
||||
return resolved
|
||||
try:
|
||||
if _supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="anthropic",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
key=key,
|
||||
):
|
||||
return True
|
||||
|
|
@ -372,17 +405,24 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
return AnthropicModelInfo._get_model_capability(model, key) is True
|
||||
|
||||
@staticmethod
|
||||
def _is_adaptive_thinking_model(model: str) -> bool:
|
||||
def _is_adaptive_thinking_model(model: str, custom_llm_provider: str) -> bool:
|
||||
"""Whether ``model`` uses adaptive thinking (``output_config.effort``).
|
||||
|
||||
The model cost map is authoritative: an explicit ``supports_adaptive_thinking``
|
||||
entry, or a ``fallback_generalizations`` rule for unknown Claude models. The
|
||||
version gate (>= 4.6, including provider-prefixed Bedrock/Vertex ids that map to
|
||||
no exact entry) lives entirely in that declarative rule, not here.
|
||||
entry resolved under ``custom_llm_provider``, or a ``fallback_generalizations``
|
||||
rule for unknown Claude models. The version gate (>= 4.6, including
|
||||
provider-prefixed Bedrock/Vertex ids that map to no exact entry) lives entirely
|
||||
in that declarative rule, not here.
|
||||
"""
|
||||
return AnthropicModelInfo._supports_model_capability(model, "supports_adaptive_thinking")
|
||||
return AnthropicModelInfo._supports_model_capability(model, "supports_adaptive_thinking", custom_llm_provider)
|
||||
|
||||
def is_effort_used(self, optional_params: Optional[dict], model: Optional[str] = None) -> bool:
|
||||
def is_effort_used(
|
||||
self,
|
||||
optional_params: Optional[dict],
|
||||
model: Optional[str] = None,
|
||||
*,
|
||||
custom_llm_provider: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if effort parameter is being used and requires a beta header.
|
||||
|
||||
|
|
@ -394,7 +434,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
return False
|
||||
|
||||
# Claude 4.6+ models use output_config as a stable API feature — no beta header needed
|
||||
if model and self._is_adaptive_thinking_model(model):
|
||||
if model and self._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
return False
|
||||
|
||||
# Check if reasoning_effort is provided for Claude Opus 4.5
|
||||
|
|
@ -475,6 +515,8 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
prompt_caching_set: bool = False,
|
||||
file_id_used: bool = False,
|
||||
mcp_server_used: bool = False,
|
||||
*,
|
||||
custom_llm_provider: str,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get list of common beta headers based on the features that are active.
|
||||
|
|
@ -487,7 +529,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
betas = []
|
||||
|
||||
# Detect features
|
||||
effort_used = self.is_effort_used(optional_params, model)
|
||||
effort_used = self.is_effort_used(optional_params, model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
if effort_used:
|
||||
betas.append(ANTHROPIC_EFFORT_BETA_HEADER) # effort-2025-11-24
|
||||
|
|
@ -643,7 +685,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
tool_search_used = self.is_tool_search_used(tools=tools)
|
||||
programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(tools=tools)
|
||||
input_examples_used = self.is_input_examples_used(tools=tools)
|
||||
effort_used = self.is_effort_used(optional_params=optional_params, model=model)
|
||||
effort_used = self.is_effort_used(optional_params=optional_params, model=model, custom_llm_provider="anthropic")
|
||||
code_execution_tool_used = self.is_code_execution_tool_used(tools=tools)
|
||||
container_with_skills_used = self.is_container_with_skills_used(optional_params=optional_params)
|
||||
user_anthropic_beta_headers = self._get_user_anthropic_beta_headers(
|
||||
|
|
|
|||
|
|
@ -32,8 +32,22 @@ from ...common_utils import (
|
|||
|
||||
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
|
||||
|
||||
DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING = (
|
||||
"Dropping adaptive `thinking`/`output_config.effort` for model=%s: the model "
|
||||
"does not support extended thinking, or max_tokens is too small to fit the "
|
||||
"minimum thinking budget."
|
||||
)
|
||||
|
||||
|
||||
class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "anthropic"
|
||||
|
||||
@property
|
||||
def _resolved_provider(self) -> str:
|
||||
return self.custom_llm_provider or "anthropic"
|
||||
|
||||
def get_supported_anthropic_messages_params(self, model: str) -> list:
|
||||
return [
|
||||
"messages",
|
||||
|
|
@ -174,7 +188,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
return headers, api_base
|
||||
|
||||
@staticmethod
|
||||
def _translate_reasoning_effort_to_anthropic(model: str, optional_params: Dict) -> None:
|
||||
def _translate_reasoning_effort_to_anthropic(model: str, optional_params: Dict, custom_llm_provider: str) -> None:
|
||||
"""Map OpenAI-style ``reasoning_effort`` to native Anthropic params.
|
||||
|
||||
Caller-supplied ``thinking`` / ``output_config`` win over the alias.
|
||||
|
|
@ -191,7 +205,11 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
return
|
||||
|
||||
try:
|
||||
mapped_thinking = AnthropicConfig._map_reasoning_effort(reasoning_effort=reasoning_effort, model=model)
|
||||
mapped_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=reasoning_effort,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except _BadRequestError as e:
|
||||
raise AnthropicError(message=str(e.message), status_code=400)
|
||||
|
||||
|
|
@ -201,7 +219,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
return
|
||||
|
||||
optional_params.setdefault("thinking", mapped_thinking)
|
||||
if AnthropicModelInfo._is_adaptive_thinking_model(model):
|
||||
if AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort)
|
||||
if mapped_effort is None:
|
||||
raise AnthropicError(
|
||||
|
|
@ -212,7 +230,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
),
|
||||
status_code=400,
|
||||
)
|
||||
gate_error = AnthropicConfig._validate_effort_for_model(model, mapped_effort)
|
||||
gate_error = AnthropicConfig._validate_effort_for_model(model, mapped_effort, custom_llm_provider)
|
||||
if gate_error is not None:
|
||||
raise AnthropicError(message=gate_error, status_code=400)
|
||||
existing_output_config = optional_params.get("output_config")
|
||||
|
|
@ -222,13 +240,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
optional_params["output_config"] = existing_output_config
|
||||
|
||||
@staticmethod
|
||||
def _translate_legacy_thinking_for_adaptive_model(model: str, optional_params: Dict) -> None:
|
||||
def _translate_legacy_thinking_for_adaptive_model(
|
||||
model: str, optional_params: Dict, custom_llm_provider: str
|
||||
) -> None:
|
||||
"""Translate legacy ``thinking.type=enabled`` to adaptive for 4.6/4.7.
|
||||
Caller-provided ``output_config.effort`` is never overridden.
|
||||
"""
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
if not AnthropicModelInfo._is_adaptive_thinking_model(model):
|
||||
if not AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
return
|
||||
thinking = optional_params.get("thinking")
|
||||
if not isinstance(thinking, dict) or thinking.get("type") != "enabled":
|
||||
|
|
@ -236,7 +256,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
|
||||
budget = int(thinking.get("budget_tokens") or 0)
|
||||
if budget >= DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET and (
|
||||
AnthropicConfig._supports_effort_level(model, "xhigh")
|
||||
AnthropicConfig._supports_effort_level(model, "xhigh", custom_llm_provider)
|
||||
):
|
||||
effort = "xhigh"
|
||||
elif budget >= DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET:
|
||||
|
|
@ -253,6 +273,108 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
existing_output_config.setdefault("effort", effort)
|
||||
optional_params["output_config"] = existing_output_config
|
||||
|
||||
@staticmethod
|
||||
def _translate_adaptive_effort_for_non_adaptive_model(
|
||||
model: str, optional_params: Dict, max_tokens: Optional[int], custom_llm_provider: str
|
||||
) -> None:
|
||||
"""Translate the 4.6+ adaptive-thinking interface (``thinking.type=adaptive``
|
||||
and/or ``output_config.effort``) down to what an older Anthropic model
|
||||
supports. Clients like Claude Code send this interface unconditionally, so
|
||||
without translation it reaches a pre-4.6 model and Anthropic rejects it with
|
||||
"This model does not support the effort parameter".
|
||||
|
||||
The reshape is silent, matching how the messages path already strips
|
||||
unsupported ``output_config`` for older models (bedrock invoke, issue
|
||||
#22797): the goal is to keep the request working, not to fail it.
|
||||
|
||||
``thinking.type=adaptive`` and ``output_config.effort`` are independent
|
||||
capabilities. Adaptive thinking needs ``supports_adaptive_thinking`` (4.6+);
|
||||
``output_config.effort`` needs ``supports_output_config``, which some
|
||||
non-adaptive models (e.g. Claude Opus 4.5) advertise on its own. So the two
|
||||
are handled separately:
|
||||
|
||||
- Adaptive-thinking models (4.6+): both are native, left untouched.
|
||||
- ``supports_output_config`` but non-adaptive (Opus 4.5): keep
|
||||
``output_config.effort`` (native), only drop the unsupported adaptive
|
||||
``thinking`` block. When adaptive thinking is being dropped and the
|
||||
effort level itself isn't supported by the model (e.g. ``xhigh``/``max``
|
||||
on Opus 4.5, which only accepts low/medium/high, while ``xhigh`` is
|
||||
Claude Code's default), fall through to the legacy translation below
|
||||
instead of forwarding a level Anthropic would reject. Effort-only
|
||||
requests are always left untouched: provider subclasses own their level
|
||||
normalization (bedrock clamps ``xhigh`` to the model's ceiling after
|
||||
this base transform runs).
|
||||
- Thinking-capable but neither (``supports_reasoning``, e.g. Haiku/Sonnet
|
||||
4.5): map effort to legacy ``thinking={type: enabled, budget_tokens}`` via
|
||||
``AnthropicConfig._map_reasoning_effort``, capped below ``max_tokens``
|
||||
(Anthropic requires ``max_tokens > budget_tokens``) and dropped when
|
||||
``max_tokens`` can't fit even the minimum budget.
|
||||
- No reasoning support: ``thinking`` is dropped.
|
||||
|
||||
For the last two, only the consumed ``effort`` key is removed from
|
||||
``output_config``; any residual (e.g. ``format``) is left for provider
|
||||
subclasses to handle.
|
||||
"""
|
||||
from litellm.exceptions import BadRequestError as _BadRequestError
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
return
|
||||
|
||||
output_config = optional_params.get("output_config")
|
||||
thinking = optional_params.get("thinking")
|
||||
effort = output_config.get("effort") if isinstance(output_config, dict) else None
|
||||
adaptive_thinking = isinstance(thinking, dict) and thinking.get("type") == "adaptive"
|
||||
if effort is None and not adaptive_thinking:
|
||||
return
|
||||
|
||||
# Models that natively accept `output_config.effort` but are not adaptive (Claude Opus 4.5).
|
||||
# Keep the native effort and only drop the adaptive `thinking` block, which these models
|
||||
# reject. Effort-only requests pass through so provider subclasses (bedrock/vertex) keep
|
||||
# owning level clamping; an adaptive request only stays here when its effort level is one
|
||||
# the model supports, otherwise it falls through to the legacy budget translation below.
|
||||
if AnthropicConfig._model_supports_effort_param(model, custom_llm_provider) and (
|
||||
not adaptive_thinking
|
||||
or AnthropicConfig._validate_effort_for_model(model, effort, custom_llm_provider) is None
|
||||
):
|
||||
if adaptive_thinking:
|
||||
optional_params.pop("thinking", None)
|
||||
return
|
||||
|
||||
supports_thinking = AnthropicModelInfo._supports_model_capability(
|
||||
model, "supports_reasoning", custom_llm_provider
|
||||
)
|
||||
try:
|
||||
legacy_thinking = (
|
||||
AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=effort or "medium",
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
if supports_thinking
|
||||
else None
|
||||
)
|
||||
except _BadRequestError as e:
|
||||
raise AnthropicError(message=str(e.message), status_code=400)
|
||||
capped_thinking = (
|
||||
AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens)
|
||||
if legacy_thinking is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if capped_thinking is not None:
|
||||
optional_params["thinking"] = capped_thinking
|
||||
else:
|
||||
verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING, model)
|
||||
optional_params.pop("thinking", None)
|
||||
|
||||
if isinstance(output_config, dict) and "effort" in output_config:
|
||||
residual = {k: v for k, v in output_config.items() if k != "effort"}
|
||||
if residual:
|
||||
optional_params["output_config"] = residual
|
||||
else:
|
||||
optional_params.pop("output_config", None)
|
||||
|
||||
def transform_anthropic_messages_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -277,11 +399,20 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
self._translate_reasoning_effort_to_anthropic(
|
||||
model=model,
|
||||
optional_params=anthropic_messages_optional_request_params,
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
)
|
||||
|
||||
self._translate_legacy_thinking_for_adaptive_model(
|
||||
model=model,
|
||||
optional_params=anthropic_messages_optional_request_params,
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
)
|
||||
|
||||
self._translate_adaptive_effort_for_non_adaptive_model(
|
||||
model=model,
|
||||
optional_params=anthropic_messages_optional_request_params,
|
||||
max_tokens=max_tokens,
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
)
|
||||
|
||||
system_param = anthropic_messages_optional_request_params.get("system")
|
||||
|
|
|
|||
|
|
@ -21,6 +21,10 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig):
|
|||
and Azure endpoint format.
|
||||
"""
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "azure_ai"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,17 @@ if TYPE_CHECKING:
|
|||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
|
||||
|
||||
|
||||
def is_azure_document_intelligence_model(model: str) -> bool:
|
||||
"""Whether an azure_ai OCR model routes to Azure Document Intelligence.
|
||||
|
||||
Azure AI exposes two OCR services on the same provider; the sub-route in the
|
||||
model name (`azure_ai/doc-intelligence/<model>`) selects Document Intelligence
|
||||
over Mistral OCR. This is the single source of truth for that routing decision.
|
||||
"""
|
||||
lowered = model.lower()
|
||||
return "doc-intelligence" in lowered or "documentintelligence" in lowered
|
||||
|
||||
|
||||
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
||||
"""
|
||||
Determine which Azure AI OCR configuration to use based on the model name.
|
||||
|
|
@ -41,7 +52,7 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
|||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
|
||||
# Check for Azure Document Intelligence models
|
||||
if "doc-intelligence" in model or "documentintelligence" in model:
|
||||
if is_azure_document_intelligence_model(model):
|
||||
verbose_logger.debug(f"Routing {model} to Azure Document Intelligence OCR config")
|
||||
return AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
|
|
|
|||
|
|
@ -50,6 +50,8 @@ _STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
|
|||
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
|
||||
)
|
||||
|
||||
SIGV4_COMPUTED_HEADERS = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
|
|
@ -1400,11 +1402,13 @@ class BaseAWSLLM:
|
|||
|
||||
# Add back all original headers (including forwarded ones) after signature calculation
|
||||
for header_name, header_value in headers.items():
|
||||
if header_value is not None:
|
||||
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
|
||||
request.headers[header_name] = header_value
|
||||
|
||||
if (
|
||||
extra_headers is not None and "Authorization" in extra_headers
|
||||
extra_headers is not None
|
||||
and "Authorization" in extra_headers
|
||||
and not extra_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
): # prevent sigv4 from overwriting the auth header
|
||||
request.headers["Authorization"] = extra_headers["Authorization"]
|
||||
prepped = request.prepare()
|
||||
|
|
@ -1527,9 +1531,15 @@ class BaseAWSLLM:
|
|||
# Add back original headers after signing. Only headers in SignedHeaders
|
||||
# are integrity-protected; forwarded headers (x-forwarded-*) must remain unsigned.
|
||||
for header_name, header_value in headers.items():
|
||||
if header_value is not None:
|
||||
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
|
||||
request_headers_dict[header_name] = header_value
|
||||
if headers is not None and "Authorization" in headers: # prevent sigv4 from overwriting the auth header
|
||||
request_headers_dict["Authorization"] = headers["Authorization"]
|
||||
incoming_authorization = next(
|
||||
(value for name, value in headers.items() if name.lower() == "authorization" and value is not None),
|
||||
None,
|
||||
)
|
||||
if incoming_authorization is not None and not incoming_authorization.startswith(
|
||||
"AWS4-HMAC-SHA256"
|
||||
): # prevent sigv4 from overwriting the auth header
|
||||
request_headers_dict["Authorization"] = incoming_authorization
|
||||
|
||||
return request_headers_dict, request.body
|
||||
|
|
|
|||
|
|
@ -5,6 +5,9 @@ from typing import Any, Dict, List, Literal, Optional, Union, cast
|
|||
|
||||
from httpx import Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
BEDROCK_MANAGED_S3_BATCH_PREFIX,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
|
@ -26,6 +29,15 @@ from litellm.types.utils import LiteLLMBatch, LlmProviders
|
|||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import CommonBatchFilesUtils
|
||||
|
||||
# Bedrock batch input files are uploaded as
|
||||
# s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see
|
||||
# BedrockFilesTransformation._get_s3_object_name). A uuid4 is always 36 hex/dash
|
||||
# characters, so it can be stripped off the end unambiguously even though the
|
||||
# model name itself may contain dashes.
|
||||
_S3_BATCH_FILE_UUID_SUFFIX_PATTERN = re.compile(
|
||||
r"-[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}\.jsonl$"
|
||||
)
|
||||
|
||||
|
||||
class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
||||
"""
|
||||
|
|
@ -40,6 +52,41 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.BEDROCK
|
||||
|
||||
@classmethod
|
||||
def _get_bare_model_name_from_s3_key(cls, object_key: str) -> Optional[str]:
|
||||
if not object_key.startswith(BEDROCK_MANAGED_S3_BATCH_PREFIX):
|
||||
return None
|
||||
model_part = object_key[len(BEDROCK_MANAGED_S3_BATCH_PREFIX) :]
|
||||
match = _S3_BATCH_FILE_UUID_SUFFIX_PATTERN.search(model_part)
|
||||
if not match or match.start() == 0:
|
||||
return None
|
||||
return model_part[: match.start()]
|
||||
|
||||
@classmethod
|
||||
def is_unmanaged_s3_batch_input_file_id(cls, input_file_id: Optional[str]) -> bool:
|
||||
"""
|
||||
Returns True if `input_file_id` is a raw s3:// Bedrock batch input file (i.e. not a
|
||||
LiteLLM-managed unified file id) whose object key embeds the model name in the
|
||||
`litellm-bedrock-files-{model}-{uuid}.jsonl` layout.
|
||||
"""
|
||||
if input_file_id is None or not input_file_id.startswith("s3://"):
|
||||
return False
|
||||
object_key = input_file_id.rsplit("/", 1)[-1]
|
||||
return cls._get_bare_model_name_from_s3_key(object_key) is not None
|
||||
|
||||
@classmethod
|
||||
def get_bare_model_name_from_s3_file(cls, input_file_id: str) -> str:
|
||||
"""
|
||||
Extracts the bare model name (e.g. "us.anthropic.claude-sonnet-4-20250514-v1-0") from
|
||||
an unmanaged batch's s3:// input file id. Note any ":" in the original model id was
|
||||
replaced with "-" at upload time, so callers must fuzzy-match against configured
|
||||
deployments rather than expect an exact string match.
|
||||
"""
|
||||
object_key = input_file_id.rsplit("/", 1)[-1]
|
||||
bare_model_name = cls._get_bare_model_name_from_s3_key(object_key)
|
||||
assert bare_model_name is not None # narrowed by is_unmanaged_s3_batch_input_file_id
|
||||
return bare_model_name
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
make_valid_bedrock_tool_name,
|
||||
)
|
||||
from litellm.llms.anthropic.chat.transformation import (
|
||||
DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING,
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT,
|
||||
AnthropicConfig,
|
||||
|
|
@ -423,6 +424,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
mapped_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=reasoning_effort,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
llm_provider="bedrock_converse",
|
||||
)
|
||||
if mapped_thinking is None:
|
||||
|
|
@ -430,7 +432,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
optional_params.pop("output_config", None)
|
||||
else:
|
||||
optional_params["thinking"] = mapped_thinking
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model, "bedrock"):
|
||||
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort)
|
||||
if mapped_effort is None:
|
||||
AnthropicConfig._raise_invalid_reasoning_effort(
|
||||
|
|
@ -465,7 +467,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
model=model,
|
||||
llm_provider="bedrock_converse",
|
||||
)
|
||||
error = AnthropicConfig._validate_effort_for_model(model=model, effort=effort)
|
||||
error = AnthropicConfig._validate_effort_for_model(model=model, effort=effort, custom_llm_provider="bedrock")
|
||||
if error is not None:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=error,
|
||||
|
|
@ -898,7 +900,28 @@ class AmazonConverseConfig(BaseConfig):
|
|||
"tool_choice": {"disable_parallel_tool_use": disable_parallel}
|
||||
}
|
||||
if param == "thinking":
|
||||
optional_params["thinking"] = value
|
||||
if (
|
||||
isinstance(value, dict)
|
||||
and value.get("type") == "adaptive"
|
||||
and not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock")
|
||||
):
|
||||
max_tokens = non_default_params.get("max_completion_tokens") or non_default_params.get("max_tokens")
|
||||
legacy_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort="medium",
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
capped = (
|
||||
AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens)
|
||||
if legacy_thinking is not None
|
||||
else None
|
||||
)
|
||||
if capped is not None:
|
||||
optional_params["thinking"] = capped
|
||||
else:
|
||||
litellm.verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, model)
|
||||
else:
|
||||
optional_params["thinking"] = value
|
||||
elif param == "reasoning_effort" and isinstance(value, str):
|
||||
self._handle_reasoning_effort_parameter(
|
||||
model=model, reasoning_effort=value, optional_params=optional_params
|
||||
|
|
@ -1279,7 +1302,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
if anthropic_output_config is not None and isinstance(anthropic_output_config, dict):
|
||||
if base_model.startswith("anthropic"):
|
||||
if litellm.drop_params is True and not AnthropicConfig._model_supports_effort_param(model):
|
||||
if litellm.drop_params is True and not AnthropicConfig._model_supports_effort_param(model, "bedrock"):
|
||||
litellm.verbose_logger.warning(
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
model,
|
||||
|
|
@ -1422,7 +1445,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if (
|
||||
isinstance(output_config, dict)
|
||||
and output_config.get("effort") is not None
|
||||
and not AnthropicConfig._is_adaptive_thinking_model(model)
|
||||
and not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock")
|
||||
):
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_EFFORT_BETA_HEADER,
|
||||
|
|
|
|||
|
|
@ -115,7 +115,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
keeps working. Non-adaptive models and models without a ceiling are
|
||||
left untouched.
|
||||
"""
|
||||
if not AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
if not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock"):
|
||||
return
|
||||
effort = params.get("reasoning_effort")
|
||||
if not isinstance(effort, str):
|
||||
|
|
@ -228,7 +228,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
or AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
if anthropic_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -269,6 +269,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
prompt_caching_set=False,
|
||||
file_id_used=self.is_file_id_used(messages),
|
||||
mcp_server_used=self.is_mcp_server_used(optional_params.get("mcp_servers")),
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
beta_set.update(auto_betas)
|
||||
|
||||
|
|
|
|||
|
|
@ -54,7 +54,9 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig):
|
|||
tool_search_used=self.is_tool_search_used(tools=optional_params.get("tools")),
|
||||
programmatic_tool_calling_used=self.is_programmatic_tool_calling_used(tools=optional_params.get("tools")),
|
||||
input_examples_used=self.is_input_examples_used(tools=optional_params.get("tools")),
|
||||
effort_used=self.is_effort_used(optional_params=optional_params, model=model),
|
||||
effort_used=self.is_effort_used(
|
||||
optional_params=optional_params, model=model, custom_llm_provider="anthropic"
|
||||
),
|
||||
user_anthropic_beta_headers=self._get_user_anthropic_beta_headers(
|
||||
anthropic_beta_header=headers.get("anthropic-beta")
|
||||
),
|
||||
|
|
|
|||
|
|
@ -77,6 +77,10 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
|
||||
DEFAULT_BEDROCK_ANTHROPIC_API_VERSION = "bedrock-2023-05-31"
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "bedrock"
|
||||
|
||||
BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS = frozenset(BedrockInvokeAnthropicMessagesRequest.__annotations__.keys())
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
|
|
@ -93,26 +97,48 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
return [{"type": "text", "text": value}]
|
||||
return [value]
|
||||
|
||||
def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict) -> None:
|
||||
"""Bedrock Invoke rejects ``role: "system"`` entries inside ``messages`` on
|
||||
some Claude aliases; Anthropic Messages carries that content in the
|
||||
top-level ``system`` field. Move any such entries into ``system`` before
|
||||
the Invoke request is built."""
|
||||
@staticmethod
|
||||
def _is_system_role_message(message: Any) -> bool:
|
||||
return isinstance(message, dict) and message.get("role") == "system"
|
||||
|
||||
def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict, model: str) -> None:
|
||||
"""Bedrock Invoke validates ``role: "system"`` entries inside ``messages``
|
||||
per model. Models carrying ``supports_mid_conversation_system`` in the
|
||||
cost map (the Opus 4.8 family) only reject a leading run ("messages.0:
|
||||
use the top-level 'system' parameter for the initial system prompt") and
|
||||
accept mid-conversation entries (e.g. Claude Code's
|
||||
``mid-conversation-system-2026-04-07`` reminders) in place, where they
|
||||
MUST stay: hoisting one mutates the ``system`` prefix and invalidates the
|
||||
prompt cache for the entire message history. Older Claude models (Opus
|
||||
4.7, Sonnet 4.6, Haiku 4.5, ...) reject the role in every position
|
||||
("role 'system' is not supported on this model"), so without the flag
|
||||
every system entry is hoisted into the top-level ``system`` field.
|
||||
Billing-header system blocks are stripped from the top-level ``system``
|
||||
field regardless of whether anything was hoisted."""
|
||||
messages = anthropic_messages_request.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
return
|
||||
system_role_messages = [m for m in messages if isinstance(m, dict) and m.get("role") == "system"]
|
||||
if not system_role_messages:
|
||||
return
|
||||
|
||||
anthropic_messages_request["messages"] = [
|
||||
m for m in messages if not (isinstance(m, dict) and m.get("role") == "system")
|
||||
]
|
||||
if _supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_mid_conversation_system",
|
||||
):
|
||||
leading_count = next(
|
||||
(i for i, m in enumerate(messages) if not self._is_system_role_message(m)),
|
||||
len(messages),
|
||||
)
|
||||
hoisted = messages[:leading_count]
|
||||
remaining = messages[leading_count:]
|
||||
else:
|
||||
hoisted = [m for m in messages if self._is_system_role_message(m)]
|
||||
remaining = [m for m in messages if not self._is_system_role_message(m)]
|
||||
if hoisted:
|
||||
anthropic_messages_request["messages"] = remaining
|
||||
system_content = [
|
||||
block
|
||||
for source in (
|
||||
anthropic_messages_request.get("system"),
|
||||
*(m.get("content") for m in system_role_messages),
|
||||
*(m.get("content") for m in hoisted),
|
||||
)
|
||||
for block in self._as_system_content_blocks(source)
|
||||
]
|
||||
|
|
@ -247,7 +273,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
Returns:
|
||||
True if the model supports extended thinking on Bedrock
|
||||
"""
|
||||
if AnthropicModelInfo._is_adaptive_thinking_model(model):
|
||||
if AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock"):
|
||||
return True
|
||||
|
||||
model_lower = model.lower()
|
||||
|
|
@ -297,7 +323,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if not self._supports_extended_thinking_on_bedrock(model):
|
||||
return False
|
||||
|
||||
is_adaptive_thinking_model = AnthropicModelInfo._is_adaptive_thinking_model(model)
|
||||
is_adaptive_thinking_model = AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock")
|
||||
|
||||
thinking = anthropic_messages_request.get("thinking")
|
||||
if isinstance(thinking, dict):
|
||||
|
|
@ -489,24 +515,43 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if self._supports_tool_search_on_bedrock(model):
|
||||
beta_set.add("tool-search-tool-2025-10-19")
|
||||
|
||||
# Bedrock-InvokeModel-supported ``context_management.edits`` types and the
|
||||
# ``anthropic-beta`` header that each one requires. ``clear_thinking_20251015``
|
||||
# is intentionally absent — it is LiteLLM-internal, consumed via
|
||||
# ``_ensure_thinking_for_clear_thinking_context_management``, and forwarding
|
||||
# the raw edit trips Bedrock's
|
||||
# ``"context_management: Extra inputs are not permitted"`` 400.
|
||||
#
|
||||
# Bedrock InvokeModel DOES support ``clear_tool_uses_20250919`` under the
|
||||
# ``context-management-2025-06-27`` beta. AWS docs:
|
||||
# https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-tool-use.md
|
||||
_BEDROCK_INVOKE_SUPPORTED_CONTEXT_MANAGEMENT_EDITS: Dict[str, str] = {
|
||||
"compact_20260112": ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value,
|
||||
"clear_tool_uses_20250919": ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _filter_context_management_for_bedrock_invoke(
|
||||
anthropic_messages_request: Dict,
|
||||
beta_set: set,
|
||||
) -> None:
|
||||
"""
|
||||
Bedrock InvokeModel accepts ``context_management`` only when it carries
|
||||
``compact_20260112`` edits paired with the ``compact-2026-01-12``
|
||||
anthropic-beta header. Other edit types (notably ``clear_thinking_20251015``,
|
||||
which Claude Code sends on every request) are LiteLLM-internal and would
|
||||
cause Bedrock to 400 with ``"context_management: Extra inputs are not
|
||||
permitted"``.
|
||||
Filter ``context_management.edits`` to the subset that Bedrock InvokeModel
|
||||
accepts and add the matching ``anthropic-beta`` header for each surviving
|
||||
edit type.
|
||||
|
||||
Filter the edits list to the supported subset, add the beta header when
|
||||
compact edits remain, and drop ``context_management`` entirely when no
|
||||
supported edits are left so the safety-net allowlist can pass it through.
|
||||
- ``compact_20260112`` -> ``compact-2026-01-12``
|
||||
- ``clear_tool_uses_20250919`` -> ``context-management-2025-06-27``
|
||||
|
||||
Ref: https://github.com/BerriAI/litellm/issues/27532
|
||||
Other edit types (notably ``clear_thinking_20251015``, which Claude Code
|
||||
sends on every request) are LiteLLM-internal: thinking is injected
|
||||
separately via ``_ensure_thinking_for_clear_thinking_context_management``,
|
||||
and forwarding the raw edit would trip Bedrock's
|
||||
``"context_management: Extra inputs are not permitted"`` 400.
|
||||
|
||||
Refs:
|
||||
* https://github.com/BerriAI/litellm/issues/27532
|
||||
* https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-tool-use.md
|
||||
"""
|
||||
cm = anthropic_messages_request.get("context_management")
|
||||
if not isinstance(cm, dict):
|
||||
|
|
@ -516,15 +561,17 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
anthropic_messages_request.pop("context_management", None)
|
||||
return
|
||||
|
||||
compact_edits = [e for e in edits if isinstance(e, dict) and e.get("type") == "compact_20260112"]
|
||||
if compact_edits:
|
||||
beta_set.add(ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value)
|
||||
anthropic_messages_request["context_management"] = {
|
||||
**cm,
|
||||
"edits": compact_edits,
|
||||
}
|
||||
else:
|
||||
supported = AmazonAnthropicClaudeMessagesConfig._BEDROCK_INVOKE_SUPPORTED_CONTEXT_MANAGEMENT_EDITS
|
||||
retained_edits = [e for e in edits if isinstance(e, dict) and e.get("type") in supported]
|
||||
if not retained_edits:
|
||||
anthropic_messages_request.pop("context_management", None)
|
||||
return
|
||||
|
||||
beta_set.update(supported[e["type"]] for e in retained_edits)
|
||||
anthropic_messages_request["context_management"] = {
|
||||
**cm,
|
||||
"edits": retained_edits,
|
||||
}
|
||||
|
||||
def _get_bedrock_invoke_anthropic_beta_headers(
|
||||
self,
|
||||
|
|
@ -553,6 +600,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
mcp_server_used=anthropic_model_info.is_mcp_server_used(
|
||||
anthropic_messages_optional_request_params.get("mcp_servers")
|
||||
),
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
beta_set.update(auto_betas)
|
||||
|
||||
|
|
@ -619,7 +667,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
path degrades ``xhigh`` -> ``max`` rather than 400-ing. Non-adaptive models
|
||||
and models without a ceiling are left untouched.
|
||||
"""
|
||||
if not AnthropicModelInfo._is_adaptive_thinking_model(model):
|
||||
if not AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock"):
|
||||
return
|
||||
effort = optional_params.get("reasoning_effort")
|
||||
if not isinstance(effort, str):
|
||||
|
|
@ -648,7 +696,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
self._normalize_system_role_messages_for_bedrock(anthropic_messages_request)
|
||||
self._normalize_system_role_messages_for_bedrock(anthropic_messages_request, model=model)
|
||||
#########################################################
|
||||
############## BEDROCK Invoke SPECIFIC TRANSFORMATION ###
|
||||
#########################################################
|
||||
|
|
@ -707,7 +755,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
or AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
if anthropic_messages_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -744,7 +792,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if (
|
||||
litellm.drop_params is True
|
||||
and "output_config" in anthropic_messages_request
|
||||
and not AnthropicConfig._model_supports_effort_param(model)
|
||||
and not AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
verbose_logger.warning(
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ Handles tiered pricing and prompt caching scenarios.
|
|||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import calculate_tiered_cost
|
||||
from litellm.types.utils import ModelInfo, Usage
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
|
@ -42,80 +43,6 @@ def _extract_token_breakdown(usage: Usage) -> TokenBreakdown:
|
|||
return TokenBreakdown(text_tokens, cached_tokens, completion_tokens, reasoning_tokens)
|
||||
|
||||
|
||||
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 * 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 * cost_per_token
|
||||
|
||||
return total_cost
|
||||
|
||||
|
||||
def _calculate_prompt_cost(
|
||||
breakdown: TokenBreakdown,
|
||||
model_info: ModelInfo,
|
||||
|
|
@ -123,12 +50,12 @@ def _calculate_prompt_cost(
|
|||
) -> float:
|
||||
"""Calculate total prompt cost including cached tokens."""
|
||||
if tiered_pricing:
|
||||
text_cost = _calculate_tiered_cost(
|
||||
text_cost = calculate_tiered_cost(
|
||||
tokens=breakdown.text_tokens,
|
||||
tiered_pricing=tiered_pricing,
|
||||
cost_key="input_cost_per_token",
|
||||
)
|
||||
cache_cost = _calculate_tiered_cost(
|
||||
cache_cost = calculate_tiered_cost(
|
||||
tokens=breakdown.cached_tokens,
|
||||
tiered_pricing=tiered_pricing,
|
||||
cost_key="cache_read_input_token_cost",
|
||||
|
|
@ -155,12 +82,12 @@ def _calculate_completion_cost(
|
|||
) -> float:
|
||||
"""Calculate total completion cost including reasoning tokens."""
|
||||
if tiered_pricing:
|
||||
completion_cost = _calculate_tiered_cost(
|
||||
completion_cost = calculate_tiered_cost(
|
||||
tokens=breakdown.completion_tokens,
|
||||
tiered_pricing=tiered_pricing,
|
||||
cost_key="output_cost_per_token",
|
||||
)
|
||||
reasoning_cost = _calculate_tiered_cost(
|
||||
reasoning_cost = calculate_tiered_cost(
|
||||
tokens=breakdown.reasoning_tokens,
|
||||
tiered_pricing=tiered_pricing,
|
||||
cost_key="output_cost_per_reasoning_token",
|
||||
|
|
|
|||
|
|
@ -181,6 +181,10 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
if key != "self" and value is not None:
|
||||
setattr(self.__class__, key, value)
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "databricks"
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
return super().get_config()
|
||||
|
|
@ -372,6 +376,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
mapped_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=reasoning_effort_value,
|
||||
model=model,
|
||||
custom_llm_provider="databricks",
|
||||
llm_provider="databricks",
|
||||
)
|
||||
if mapped_thinking is None:
|
||||
|
|
@ -379,7 +384,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
optional_params.pop("output_config", None)
|
||||
else:
|
||||
optional_params["thinking"] = mapped_thinking
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model, "databricks"):
|
||||
mapped_effort: Optional[str] = None
|
||||
if isinstance(reasoning_effort_value, str):
|
||||
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort_value)
|
||||
|
|
|
|||
|
|
@ -35,7 +35,9 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
|
|||
Map OpenAI params to DeepSeek params.
|
||||
|
||||
Handles `thinking` and `reasoning_effort` parameters for DeepSeek reasoner models.
|
||||
DeepSeek only supports `{"type": "enabled"}` - no budget_tokens like Anthropic.
|
||||
DeepSeek supports `{"type": "enabled"}` and `{"type": "disabled"}` - no budget_tokens
|
||||
like Anthropic. `reasoning_effort="none"` is the OpenAI-style way to ask for thinking
|
||||
off, so it maps to `{"type": "disabled"}`; any other effort keeps thinking on.
|
||||
|
||||
Reference: https://api-docs.deepseek.com/guides/thinking_mode
|
||||
"""
|
||||
|
|
@ -47,15 +49,13 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
|
|||
thinking_value = optional_params.pop("thinking", None)
|
||||
reasoning_effort = optional_params.pop("reasoning_effort", None)
|
||||
|
||||
# Handle thinking parameter - only accept {"type": "enabled"}
|
||||
if thinking_value is not None:
|
||||
if isinstance(thinking_value, dict) and thinking_value.get("type") == "enabled":
|
||||
# DeepSeek only accepts {"type": "enabled"}, ignore budget_tokens
|
||||
optional_params["thinking"] = {"type": "enabled"}
|
||||
# Handle thinking parameter - accept both enabled and disabled, ignore budget_tokens
|
||||
if isinstance(thinking_value, dict) and thinking_value.get("type") in ("enabled", "disabled"):
|
||||
optional_params["thinking"] = {"type": thinking_value["type"]}
|
||||
|
||||
# Handle reasoning_effort - map to thinking enabled
|
||||
elif reasoning_effort is not None and reasoning_effort != "none":
|
||||
optional_params["thinking"] = {"type": "enabled"}
|
||||
# Otherwise fall back to reasoning_effort: "none" disables, anything else enables
|
||||
elif reasoning_effort is not None:
|
||||
optional_params["thinking"] = {"type": "disabled" if reasoning_effort == "none" else "enabled"}
|
||||
|
||||
return optional_params
|
||||
|
||||
|
|
|
|||
|
|
@ -25,6 +25,10 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig):
|
|||
super().__init__()
|
||||
self.authenticator = Authenticator()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "github_copilot"
|
||||
|
||||
def handles_web_search_natively(self) -> bool:
|
||||
"""
|
||||
Copilot's /v1/messages endpoint does not execute ``web_search`` tools, so
|
||||
|
|
|
|||
|
|
@ -91,7 +91,7 @@ def create_config_class(provider: SimpleProviderConfig):
|
|||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""Get supported OpenAI params, excluding tool-related params for models
|
||||
that don't support function calling."""
|
||||
from litellm.utils import supports_function_calling
|
||||
from litellm.utils import supports_function_calling, supports_reasoning
|
||||
|
||||
supported_params = super().get_supported_openai_params(model=model)
|
||||
|
||||
|
|
@ -113,6 +113,10 @@ def create_config_class(provider: SimpleProviderConfig):
|
|||
f"function calling — removed tool-related params from supported params."
|
||||
)
|
||||
|
||||
_supports_reasoning = supports_reasoning(model=model, custom_llm_provider=provider.slug)
|
||||
if _supports_reasoning and "reasoning_effort" not in supported_params:
|
||||
supported_params.append("reasoning_effort")
|
||||
|
||||
return supported_params
|
||||
|
||||
def map_openai_params(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
from typing import Any, Optional
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.llms.openai_like.json_loader import SimpleProviderConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
|
||||
|
||||
|
|
@ -67,3 +70,69 @@ class OpenAILikeAnthropicMessagesConfig(AnthropicMessagesConfig):
|
|||
if base.endswith("/v1"):
|
||||
base = base[: -len("/v1")]
|
||||
return f"{base}/v1/messages"
|
||||
|
||||
|
||||
class JSONProviderAnthropicMessagesConfig(OpenAILikeAnthropicMessagesConfig):
|
||||
"""
|
||||
Provider-level native Anthropic Messages passthrough for JSON-configured
|
||||
OpenAI-compatible providers whose ``supported_endpoints`` in providers.json
|
||||
includes ``"/v1/messages"``. Resolves the api key and api base from the
|
||||
provider's configured env vars, then forwards the Anthropic payload
|
||||
untranslated like ``OpenAILikeAnthropicMessagesConfig``.
|
||||
"""
|
||||
|
||||
def __init__(self, provider: SimpleProviderConfig):
|
||||
super().__init__()
|
||||
self._provider = provider
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return self._provider.slug
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def _resolve_api_key(self, api_key: Optional[str]) -> Optional[str]:
|
||||
return api_key or get_secret_str(self._provider.api_key_env) or litellm.api_key
|
||||
|
||||
def _resolve_api_base(self, api_base: Optional[str]) -> str:
|
||||
env_api_base = get_secret_str(self._provider.api_base_env) if self._provider.api_base_env else None
|
||||
return api_base or env_api_base or self._provider.base_url
|
||||
|
||||
def validate_anthropic_messages_environment(
|
||||
self,
|
||||
headers: dict[str, str],
|
||||
model: str,
|
||||
messages: list[Any],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> tuple[dict[str, str], Optional[str]]:
|
||||
return super().validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
api_key=self._resolve_api_key(api_key),
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
return super().get_complete_url(
|
||||
api_base=self._resolve_api_base(api_base),
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -168,6 +168,13 @@
|
|||
},
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
|
||||
},
|
||||
"meta": {
|
||||
"base_url": "https://api.meta.ai/v1",
|
||||
"api_key_env": "META_API_KEY",
|
||||
"api_base_env": "META_API_BASE",
|
||||
"base_class": "openai_gpt",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/messages"]
|
||||
},
|
||||
"pinstripes": {
|
||||
"base_url": "https://pinstripes.io/v1",
|
||||
"api_key_env": "PINSTRIPES_API_KEY",
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ from ..output_params_utils import sanitize_vertex_anthropic_output_params
|
|||
|
||||
|
||||
class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase):
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "vertex_ai"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ def _model_accepts_output_config_effort(model: str) -> bool:
|
|||
"""
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
return AnthropicConfig._model_supports_effort_param(model)
|
||||
return AnthropicConfig._model_supports_effort_param(model, "vertex_ai")
|
||||
|
||||
|
||||
def sanitize_vertex_anthropic_output_params(data: dict, model: str) -> None:
|
||||
|
|
|
|||
|
|
@ -112,6 +112,7 @@ class VertexAIAnthropicConfig(AnthropicConfig):
|
|||
prompt_caching_set=self.is_cache_control_set(messages),
|
||||
file_id_used=self.is_file_id_used(messages),
|
||||
mcp_server_used=self.is_mcp_server_used(optional_params.get("mcp_servers")),
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
beta_set = set(auto_betas)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -83,10 +83,20 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
token_url: Optional[str] = None
|
||||
registration_url: Optional[str] = None
|
||||
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
|
||||
# Token Exchange (OBO) fields — RFC 8693. ``audience`` is named for the RFC's
|
||||
# request parameter (token-exchange only); RFC 8707 resource indicators are a
|
||||
# separate concept named ``resource`` in the v2 egress types. A null
|
||||
# ``subject_token_type`` means DEFAULT_SUBJECT_TOKEN_TYPE (litellm.types.mcp),
|
||||
# applied at the egress build sites.
|
||||
token_exchange_endpoint: Optional[str] = None
|
||||
audience: Optional[str] = None
|
||||
subject_token_type: Optional[str] = None
|
||||
token_exchange_profile: Optional[str] = None
|
||||
allow_all_keys: bool = False
|
||||
available_on_public_internet: bool = True
|
||||
delegate_auth_to_upstream: bool = False
|
||||
oauth_passthrough: bool = False
|
||||
dcr_bridge: Optional[bool] = None
|
||||
is_byok: bool = False
|
||||
byok_description: List[str] = Field(default_factory=list)
|
||||
byok_api_key_help_url: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -17,6 +17,9 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import (
|
||||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
|
|
@ -83,6 +86,8 @@ def _prepare_ocr_request(
|
|||
if doc_type not in ["document_url", "image_url"]:
|
||||
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
|
||||
|
||||
caller_supplied_api_base = api_base is not None
|
||||
|
||||
(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
|
|
@ -95,9 +100,14 @@ def _prepare_ocr_request(
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
suppress_dynamic_api_base = (
|
||||
not caller_supplied_api_base
|
||||
and custom_llm_provider == "azure_ai"
|
||||
and is_azure_document_intelligence_model(model)
|
||||
)
|
||||
if dynamic_api_key:
|
||||
api_key = dynamic_api_key
|
||||
if dynamic_api_base:
|
||||
if dynamic_api_base and not suppress_dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
ocr_provider_config = ProviderConfigManager.get_provider_ocr_config(
|
||||
|
|
@ -191,8 +201,7 @@ def _rust_bridge_api_base(
|
|||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
model = prepared_request.model.lower()
|
||||
if "doc-intelligence" in model or "documentintelligence" in model:
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
|||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
|
@ -35,8 +36,6 @@ if TYPE_CHECKING:
|
|||
# RFC 8693 grant type constant
|
||||
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
|
||||
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token"
|
||||
|
||||
|
||||
class TokenExchangeHandler:
|
||||
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.
|
||||
|
|
|
|||
|
|
@ -1,12 +1,23 @@
|
|||
import re
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List, Optional, Set, Tuple, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
from starlette.types import Scope
|
||||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
|
||||
BridgeEnvelopeAdmitted,
|
||||
BridgeEnvelopeInvalid,
|
||||
NotBridgeEnvelope,
|
||||
envelope_keys_from_master_key,
|
||||
is_bridge_envelope_shaped,
|
||||
resolve_bridge_envelope,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
UI_TEAM_ID,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -17,12 +28,17 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_run_centralized_common_checks,
|
||||
user_api_key_auth,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
from litellm.repositories.table_repositories import (
|
||||
AgentsRepository,
|
||||
MCPServerRepository,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[List[str]] = None) -> Optional[List[str]]:
|
||||
|
|
@ -220,6 +236,35 @@ class MCPRequestHandler:
|
|||
# when EVERY target is auth_type=oauth2 with delegate_auth_to_upstream
|
||||
# set; fails closed otherwise.
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif MCPRequestHandler._target_servers_are_true_passthrough(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif (
|
||||
(
|
||||
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
)
|
||||
is not None
|
||||
and oauth2_headers
|
||||
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
|
||||
):
|
||||
# A single DCR-bridge oauth_delegate target carrying an envelope-shaped
|
||||
# Authorization: open the envelope, admit under its recovered identity, and
|
||||
# inject the inner upstream token for egress. A non-envelope bearer on the same
|
||||
# server is NOT admitted here — it falls through to the oauth2 arm, which 401s.
|
||||
validated_user_api_key_auth, mcp_server_auth_headers = await MCPRequestHandler._admit_dcr_bridge_delegate(
|
||||
server=bridge_delegate_target,
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request=request,
|
||||
route=request_route,
|
||||
)
|
||||
elif oauth2_headers:
|
||||
# Authorization on a non-delegated server: the bearer must be a real
|
||||
# LiteLLM credential, so a failed validation is a genuine 401/403 and
|
||||
|
|
@ -399,6 +444,274 @@ class MCPRequestHandler:
|
|||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _target_servers_are_true_passthrough(
|
||||
path: str, mcp_servers: Optional[list[str]], client_ip: Optional[str]
|
||||
) -> bool:
|
||||
"""
|
||||
True only when EVERY MCP server the request targets is ``auth_type == true_passthrough``.
|
||||
Fails closed when any target does not opt in or cannot be resolved.
|
||||
|
||||
Used by :meth:`process_mcp_request` to skip LiteLLM admission auth entirely: the gateway is a
|
||||
transparent proxy and the caller's ``Authorization`` is an upstream token, never a LiteLLM key.
|
||||
Mirrors :meth:`_target_servers_delegate_auth_to_upstream`; a mixed-target request keeps normal auth.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
target_names = MCPRequestHandler._resolve_target_server_names(path=path, mcp_servers_header=mcp_servers)
|
||||
if not target_names:
|
||||
return False
|
||||
|
||||
for name in target_names:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_name(name, client_ip=client_ip)
|
||||
if server is None or server.auth_type != MCPAuth.true_passthrough:
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _single_dcr_bridge_delegate_target(
|
||||
path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str]
|
||||
) -> Optional[MCPServer]:
|
||||
"""The one DCR-bridge ``oauth_delegate`` server this request targets, or ``None``.
|
||||
|
||||
Returns the server only when EXACTLY ONE target resolves and it is both
|
||||
``is_oauth_delegate`` and ``is_dcr_bridge``. Fails closed (``None``) on a
|
||||
multi-target request, an unresolved target, or a non-matching server, so the
|
||||
envelope admission arm never fires for an aggregate scope or a server that did not
|
||||
opt into the bridge. Mirrors :meth:`_target_servers_are_true_passthrough`.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
target_names = MCPRequestHandler._resolve_target_server_names(path=path, mcp_servers_header=mcp_servers)
|
||||
if len(target_names) != 1:
|
||||
return None
|
||||
server = global_mcp_server_manager.get_mcp_server_by_name(target_names[0], client_ip=client_ip)
|
||||
if server is None or not server.is_oauth_delegate or not server.is_dcr_bridge:
|
||||
return None
|
||||
# Egress resolves the injected per-server token only by alias / server_name; a server with
|
||||
# neither cannot receive the forwarded token, so fail closed rather than admit-and-drop.
|
||||
if not (server.server_name or server.alias):
|
||||
return None
|
||||
return server
|
||||
|
||||
@staticmethod
|
||||
async def _admit_dcr_bridge_delegate(
|
||||
server: MCPServer,
|
||||
authorization_value: str,
|
||||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
|
||||
request: Request,
|
||||
route: str,
|
||||
) -> Tuple[UserAPIKeyAuth, Optional[Dict[str, Dict[str, str]]]]:
|
||||
"""Open the bridge envelope and admit the caller under the live key it references.
|
||||
|
||||
The envelope's signature proves the user authenticated when it was minted, but
|
||||
authorization is resolved fresh here rather than trusted from the envelope: the
|
||||
sealed ``key_hash`` reloads the current ``UserAPIKeyAuth`` record, and the admitted
|
||||
identity then runs through the standard pipeline's centralized policy gate, so the
|
||||
key's present restrictions and revocation state gate the request instead of a
|
||||
snapshot frozen at mint time. The inner upstream token is injected under the
|
||||
server's per-server auth-header key so egress forwards it via the
|
||||
``PassthroughConfig`` override; the envelope ``Authorization`` the leak-defense
|
||||
strips never reaches the upstream. A new headers dict is returned rather than
|
||||
mutating the input. Fails closed with a 401 on an invalid or expired envelope, or
|
||||
when the referenced key is missing, blocked, or expired, its owner is
|
||||
SCIM-deactivated, or the centralized policy gate rejects it (blocked team or
|
||||
project, org or budget limits).
|
||||
|
||||
The sealed token is keyed alias-first, matching the order egress resolves
|
||||
(``lookup_mcp_server_auth_in_headers`` tries ``alias`` before ``server_name``). Keying
|
||||
under ``server_name`` would leave a caller-supplied ``x-mcp-{alias}-authorization`` at the
|
||||
higher-priority alias slot, pairing the admitted identity with an attacker's upstream
|
||||
credential; the alias-keyed injection overwrites any such caller value.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import master_key
|
||||
|
||||
if not master_key:
|
||||
raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set")
|
||||
|
||||
await MCPRequestHandler._run_pre_db_read_auth_checks(request=request, route=route)
|
||||
|
||||
keys = envelope_keys_from_master_key(master_key)
|
||||
result = resolve_bridge_envelope(authorization_value, keys, datetime.now(timezone.utc), server.server_id)
|
||||
match result:
|
||||
case BridgeEnvelopeAdmitted():
|
||||
header_key = server.alias or server.server_name
|
||||
if header_key is None:
|
||||
raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name")
|
||||
admitted = await MCPRequestHandler._reload_admitted_key(result.identity.key_hash)
|
||||
await MCPRequestHandler._enforce_admitted_live_policy(admitted=admitted, request=request, route=route)
|
||||
injected = {header_key: {"Authorization": result.upstream_authorization.get_secret_value()}}
|
||||
new_headers = {**(mcp_server_auth_headers or {}), **injected}
|
||||
return admitted, new_headers
|
||||
case BridgeEnvelopeInvalid() | NotBridgeEnvelope():
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
case _:
|
||||
assert_never(result)
|
||||
|
||||
@staticmethod
|
||||
async def _run_pre_db_read_auth_checks(request: Request, route: str) -> None:
|
||||
"""Run the proxy-wide gates ``user_api_key_auth`` applies before any key lookup: the
|
||||
request-size and body-safety limits, the IP allowlist, and the ``general_settings``
|
||||
route allowlist. The envelope arm bypasses ``user_api_key_auth`` (it opens the envelope
|
||||
and reloads the identity itself), so without this a caller blocked by IP or hitting a
|
||||
proxy route the allowlist forbids would be admitted through an envelope where the same
|
||||
principal presented on the normal MCP admission path would be rejected. Runs before the
|
||||
envelope crypto so a disallowed caller is turned away before any work, mirroring the
|
||||
standard pipeline's pre-DB ordering. Violations raise the gate's own status (an IP or
|
||||
route block is a 403, an oversized body its own limit error)."""
|
||||
from litellm.proxy.auth.auth_utils import pre_db_read_auth_checks
|
||||
|
||||
await pre_db_read_auth_checks(
|
||||
request=request,
|
||||
request_data=await _read_request_body(request=request),
|
||||
route=route,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth:
|
||||
"""Reload the live key record an admitted envelope references and re-check live policy.
|
||||
|
||||
Resolving the current ``UserAPIKeyAuth`` (cache first, then DB) is what stops the
|
||||
envelope from carrying frozen authority: the key's present team/org/object-permission
|
||||
restrictions ride on the returned object, and a key that has since been deleted,
|
||||
blocked, or expired fails closed with a 401 here rather than being admitted as an
|
||||
unrestricted identity. ``get_key_object`` raises for a hash with no key row; a
|
||||
blocked or expired row is rejected explicitly because ``get_key_object`` resolves a
|
||||
row without applying those checks (the main ``user_api_key_auth`` pipeline enforces
|
||||
them downstream, which this admission path bypasses). The owner's SCIM state is the
|
||||
other builder-inline check mirrored here, so IdP offboarding revokes every envelope
|
||||
minted under the user's keys rather than leaving them live until expiry. Team,
|
||||
project, org, and budget state are NOT re-checked here; the caller runs the admitted
|
||||
identity through ``_enforce_admitted_live_policy`` for those.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_key_object
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Server misconfigured: no database connection")
|
||||
try:
|
||||
key_object = await get_key_object(
|
||||
hashed_token=key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
except Exception as e: # noqa: BLE001 # a DB outage during reload is a retryable 503, not an opaque 500
|
||||
MCPRequestHandler._raise_503_if_db_unavailable(e)
|
||||
raise
|
||||
if not MCPRequestHandler._admitted_key_is_active(key_object):
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
await MCPRequestHandler._reject_if_admitted_owner_scim_deactivated(key_object)
|
||||
return key_object
|
||||
|
||||
@staticmethod
|
||||
def _raise_503_if_db_unavailable(e: Exception) -> None:
|
||||
"""Raise a retryable 503 when ``e`` means the auth database is unreachable, else return so the
|
||||
caller applies its own fail-closed mapping. A DB outage must not masquerade as an auth failure
|
||||
(401) or surface as an opaque 500; the caller retries. Mirrors ``UserAPIKeyAuthExceptionHandler``,
|
||||
which renders a service-unavailable database error as 503 on the standard pipeline."""
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="Service Unavailable, the authentication database is temporarily unreachable. Please retry shortly.",
|
||||
) from None
|
||||
|
||||
@staticmethod
|
||||
async def _reject_if_admitted_owner_scim_deactivated(key_object: UserAPIKeyAuth) -> None:
|
||||
"""Fail closed with a 401 when the key's owning user was deactivated via SCIM.
|
||||
|
||||
The standard pipeline enforces this inline in ``_user_api_key_auth_builder`` rather
|
||||
than in ``common_checks``, so the centralized policy gate does not cover it; without
|
||||
this mirror, IdP offboarding would leave the user's already-minted envelopes live
|
||||
until expiry. A failed user lookup skips the gate (fail-open), matching the builder:
|
||||
this is the one deliberately fail-open check in an otherwise fail-closed arm, so a
|
||||
transient DB outage during this lookup admits the request rather than rejecting it,
|
||||
keeping parity with how the standard pipeline treats the same lookup failure."""
|
||||
if key_object.user_id is None:
|
||||
return
|
||||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
try:
|
||||
user_object = await get_user_object(
|
||||
user_id=key_object.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # mirror the builder's fail-open user lookup; DB errors are of any type
|
||||
verbose_logger.debug(f"bridge admission: user lookup failed, skipping SCIM gate: {e}")
|
||||
user_object = None
|
||||
if user_object is None or not isinstance(user_object.metadata, dict):
|
||||
return
|
||||
if user_object.metadata.get("scim_active") is False:
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
|
||||
@staticmethod
|
||||
async def _enforce_admitted_live_policy(admitted: UserAPIKeyAuth, request: Request, route: str) -> None:
|
||||
"""Run the standard pipeline's authorization checks over the admitted identity.
|
||||
|
||||
Mirrors the ``user_api_key_auth`` wrapper between the builder and its return: clear the
|
||||
request-scoped ``budget_reservation`` on the reloaded identity, run the route gate
|
||||
(``RouteChecks.should_call_route``) to enforce the identity's ``allowed_routes`` and any
|
||||
disabled/admin-only route, then run ``_run_centralized_common_checks`` (the same gate every
|
||||
builder path funnels through) for team-block, project-block, org, and budget. The route gate
|
||||
closes a bypass: a key barred from MCP routes could otherwise mint an envelope at the token
|
||||
endpoint (not itself an MCP route) and replay it against MCP, because the centralized checks
|
||||
treat MCP as an inference route and never re-check ``allowed_routes``.
|
||||
|
||||
Failures surface with the status the standard pipeline would give them, mirroring
|
||||
``UserAPIKeyAuthExceptionHandler``: a disallowed route is the route gate's own 403, an
|
||||
over-budget identity is a 429, a sub-check that raised its own ``HTTPException``/
|
||||
``ProxyException`` keeps that status, a transient database outage is a retryable 503, and
|
||||
only a genuinely unresolvable failure (a blocked team/project raises a bare ``Exception``,
|
||||
same as the standard pipeline's fallback) becomes the fail-closed 401. Collapsing every
|
||||
failure to 401 was misleading: it told an over-budget but validly-authenticated caller their
|
||||
credential was invalid, which on a DCR client reads as broken auth and can trigger a
|
||||
pointless re-authorize loop that cannot fix a budget problem, and it masked a DB outage as an
|
||||
auth error."""
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
admitted.budget_reservation = None
|
||||
try:
|
||||
RouteChecks.should_call_route(route=route, valid_token=admitted, request=request)
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=admitted,
|
||||
request=request,
|
||||
request_data=await _read_request_body(request=request),
|
||||
route=route,
|
||||
)
|
||||
except (HTTPException, ProxyException):
|
||||
raise
|
||||
except litellm.BudgetExceededError as e:
|
||||
raise HTTPException(status_code=getattr(e, "status_code", 429), detail=str(e)) from None
|
||||
except Exception as e: # noqa: BLE001 # untyped gate failure: retryable 503 for a DB outage, else fail closed 401
|
||||
MCPRequestHandler._raise_503_if_db_unavailable(e)
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
|
||||
@staticmethod
|
||||
def _admitted_key_is_active(key_object: UserAPIKeyAuth) -> bool:
|
||||
"""False when the referenced key is blocked or past its expiry, so a revoked key
|
||||
cannot be admitted through its still-unexpired envelope. Mirrors the active-key gate
|
||||
the bridge token endpoint applies at mint time."""
|
||||
if key_object.blocked is True:
|
||||
return False
|
||||
expires = key_object.expires
|
||||
if expires is None:
|
||||
return True
|
||||
expiry = expires if isinstance(expires, datetime) else datetime.fromisoformat(expires)
|
||||
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
return expiry >= datetime.now(timezone.utc)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_target_server_names(path: str, mcp_servers_header: Optional[List[str]]) -> List[str]:
|
||||
"""
|
||||
|
|
@ -558,10 +871,19 @@ class MCPRequestHandler:
|
|||
|
||||
ASGI headers are in format: List[List[bytes, bytes]]
|
||||
We need to convert them to the format Headers expects.
|
||||
|
||||
Collapsing the ASGI list into a dict keeps the last value for a duplicated
|
||||
header name, so a request carrying more than one ``Authorization`` is
|
||||
rejected first: for the client-forwarded token modes the gateway relays the
|
||||
caller's ``Authorization`` upstream, so a duplicate would make which token is
|
||||
forwarded ambiguous (and diverge from what admission inspected). Multiple
|
||||
``Authorization`` headers is malformed for bearer auth anyway (RFC 9110: not
|
||||
a comma-combinable field), so fail closed with a 400.
|
||||
"""
|
||||
raw_headers = scope.get("headers", [])
|
||||
MCPRequestHandler._reject_duplicate_authorization(raw_headers)
|
||||
try:
|
||||
# ASGI headers are list of [name: bytes, value: bytes] pairs
|
||||
raw_headers = scope.get("headers", [])
|
||||
# Convert bytes to strings and create dict for Headers constructor
|
||||
headers_dict = {name.decode("latin-1"): value.decode("latin-1") for name, value in raw_headers}
|
||||
return Headers(headers_dict)
|
||||
|
|
@ -570,6 +892,26 @@ class MCPRequestHandler:
|
|||
# Return empty Headers object with empty dict
|
||||
return Headers({})
|
||||
|
||||
@staticmethod
|
||||
def _reject_duplicate_authorization(raw_headers: object) -> None:
|
||||
"""Raise 400 when the raw ASGI headers carry more than one ``Authorization`` header."""
|
||||
if not isinstance(raw_headers, (list, tuple)):
|
||||
return
|
||||
count = 0
|
||||
for entry in raw_headers:
|
||||
if not isinstance(entry, (list, tuple)) or len(entry) < 1:
|
||||
continue
|
||||
name = entry[0]
|
||||
if isinstance(name, (bytes, bytearray)) and bytes(name).lower() == b"authorization":
|
||||
count += 1
|
||||
elif isinstance(name, str) and name.lower() == "authorization":
|
||||
count += 1
|
||||
if count > 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Multiple Authorization headers are not allowed",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def get_allowed_mcp_servers(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import binascii
|
|||
import hashlib
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Set, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Dict, Iterable, List, Optional, Set, Union, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -46,6 +46,64 @@ from litellm.types.mcp import MCPCredentials
|
|||
if TYPE_CHECKING:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
|
||||
{
|
||||
"authorization_url",
|
||||
"token_url",
|
||||
"registration_url",
|
||||
"oauth2_flow",
|
||||
"dcr_bridge",
|
||||
"token_exchange_endpoint",
|
||||
"audience",
|
||||
"subject_token_type",
|
||||
"token_exchange_profile",
|
||||
}
|
||||
)
|
||||
|
||||
# Token-exchange settings with dedicated columns that also exist on
|
||||
# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
|
||||
# columns). Every write lifts blob values into the columns and strips them from
|
||||
# the stored blob, so the read-time ``column or blob`` fallback only serves rows
|
||||
# the current code has never written — a cleared column can then never be
|
||||
# silently resurrected by a stale blob copy. These keys are stored plaintext
|
||||
# (endpoints/identifiers, not secrets), so values lift as-is.
|
||||
_TOKEN_EXCHANGE_COLUMN_FIELDS: frozenset = frozenset(
|
||||
{
|
||||
"token_exchange_endpoint",
|
||||
"audience",
|
||||
"subject_token_type",
|
||||
"token_exchange_profile",
|
||||
}
|
||||
)
|
||||
|
||||
# The client-forwarded token modes share one stored-credential shape: the admin-declared upstream
|
||||
# OAuth app (client_id/client_secret) plus the same authorize relay, and neither mints anything the
|
||||
# gateway keeps. So a switch WITHIN this class must preserve the stored app, unlike a cross-class
|
||||
# switch (e.g. an oauth2 row whose client may be DCR-minted and is not reusable elsewhere).
|
||||
_CLIENT_FORWARDED_AUTH_TYPES: frozenset = frozenset({"true_passthrough", "oauth_delegate"})
|
||||
|
||||
# Minted token material that must never survive a client rotation on a persisted row.
|
||||
_MINTED_TOKEN_CREDENTIAL_FIELDS: frozenset = frozenset({"access_token", "refresh_token", "expires_in"})
|
||||
|
||||
|
||||
def _credential_auth_class(auth_type: Optional[str]) -> Optional[str]:
|
||||
"""Collapse the client-forwarded modes to one credential class; every other auth_type is its own
|
||||
class. Used so credential handling keys off whether the stored-credential shape actually changed,
|
||||
not off a raw auth_type inequality that treats true_passthrough<->oauth_delegate as a full reset."""
|
||||
if auth_type in _CLIENT_FORWARDED_AUTH_TYPES:
|
||||
return "client_forwarded"
|
||||
return auth_type
|
||||
|
||||
|
||||
def _drop_stale_minted_on_client_rotation(merged: Dict[str, Any], new_creds: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""When the update rotates the client, drop stale minted token keys it did not itself set, so an old
|
||||
app's access/refresh token never rides forward under the new client. A no-op when no client key changed."""
|
||||
if "client_id" not in new_creds and "client_secret" not in new_creds:
|
||||
return merged
|
||||
return {
|
||||
key: value for key, value in merged.items() if key not in _MINTED_TOKEN_CREDENTIAL_FIELDS or key in new_creds
|
||||
}
|
||||
|
||||
|
||||
def _is_global_env_var_scope(scope: Any) -> bool:
|
||||
"""``scope="user"`` entries are placeholders the user fills in; everything
|
||||
|
|
@ -241,6 +299,14 @@ def _prepare_mcp_server_data(
|
|||
# Handle credentials serialization
|
||||
credentials = data_dict.get("credentials")
|
||||
if credentials is not None:
|
||||
# Lift legacy blob-shaped token-exchange settings into their dedicated
|
||||
# columns (an explicit top-level value wins, including an explicit
|
||||
# null) and strip them from the blob so it never seeds the read-time
|
||||
# fallback for rows written by current code.
|
||||
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
|
||||
blob_value = credentials.pop(te_field, None)
|
||||
if blob_value is not None and te_field not in data_dict:
|
||||
data_dict[te_field] = blob_value
|
||||
data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key())
|
||||
data_dict["credentials"] = safe_dumps(data_dict["credentials"])
|
||||
|
||||
|
|
@ -521,7 +587,11 @@ async def delete_mcp_server_from_virtualkey():
|
|||
pass
|
||||
|
||||
|
||||
async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Optional[LiteLLM_MCPServerTable]:
|
||||
async def delete_mcp_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> Optional[LiteLLM_MCPServerTable]:
|
||||
"""
|
||||
Delete the mcp server from the db by server_id
|
||||
|
||||
|
|
@ -532,6 +602,12 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
caller-visible error. Each table is cleaned independently so a failure on one
|
||||
still attempts the other.
|
||||
|
||||
Each enumerated credential row's user also gets their cached per-user token
|
||||
invalidated (legacy cache + v2 store, via invalidate_token_cache, defaulting
|
||||
to the manager's shared invalidation): the caches are keyed by
|
||||
(user_id, server_id), so without this a re-created server reusing the same
|
||||
server_id would serve tokens minted for the deleted server until TTL.
|
||||
|
||||
Returns the deleted mcp server record if it exists, otherwise None
|
||||
"""
|
||||
deleted_server = await MCPServerRepository(prisma_client).table.delete(
|
||||
|
|
@ -540,6 +616,18 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
},
|
||||
)
|
||||
if deleted_server is not None:
|
||||
credential_user_ids: List[str] = []
|
||||
try:
|
||||
credential_rows = await prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
where={"server_id": server_id}
|
||||
)
|
||||
credential_user_ids = [row.user_id for row in credential_rows]
|
||||
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s deleted but per-user credential enumeration failed; cached tokens expire by TTL: %s",
|
||||
server_id,
|
||||
e,
|
||||
)
|
||||
for model, label in (
|
||||
(prisma_client.db.litellm_mcpusercredentials, "credential"),
|
||||
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
|
||||
|
|
@ -554,6 +642,15 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
label,
|
||||
e,
|
||||
)
|
||||
if credential_user_ids:
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
for user_id in credential_user_ids:
|
||||
await invalidate_token_cache(user_id, server_id)
|
||||
return deleted_server
|
||||
|
||||
|
||||
|
|
@ -603,29 +700,54 @@ async def update_mcp_server(
|
|||
# Pre-fetch existing record once if we need it for auth_type or credential logic
|
||||
existing = None
|
||||
has_credentials = "credentials" in data_dict and data_dict["credentials"] is not None
|
||||
if data.auth_type or has_credentials:
|
||||
# An explicit token-exchange column write (set or clear) also migrates the
|
||||
# legacy blob copies below, so the existing row is needed for those updates.
|
||||
explicit_te_write = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
|
||||
if data.auth_type or has_credentials or explicit_te_write:
|
||||
existing = await MCPServerRepository(prisma_client).table.find_unique(where={"server_id": data.server_id})
|
||||
|
||||
# Clear stale credentials when auth_type changes but no new credentials provided
|
||||
if (
|
||||
auth_type_changed = bool(
|
||||
data.auth_type
|
||||
and "credentials" not in data_dict
|
||||
and existing
|
||||
and existing.auth_type is not None
|
||||
and existing.auth_type != data.auth_type
|
||||
):
|
||||
and _credential_auth_class(existing.auth_type) != _credential_auth_class(data.auth_type)
|
||||
)
|
||||
|
||||
# Clear stale credentials when auth_type changes but no new credentials provided
|
||||
if auth_type_changed and "credentials" not in data_dict:
|
||||
data_dict["credentials"] = None
|
||||
|
||||
if auth_type_changed:
|
||||
data_dict.update({field: None for field in _AUTH_FLOW_SCOPED_FIELDS if field not in data_dict})
|
||||
|
||||
# An explicit column write that does not touch credentials must still migrate
|
||||
# the row's legacy blob copies: lift values for columns the caller left
|
||||
# untouched, strip every copy from the blob. Without this, clearing a column
|
||||
# (e.g. to re-enable RFC 9728/8414 discovery) would leave the blob copy in
|
||||
# place, and the next credentials update's migrate-on-write would silently
|
||||
# repopulate the column the admin just cleared. (When credentials ARE in the
|
||||
# update, the merge below performs the same migration.)
|
||||
if explicit_te_write and "credentials" not in data_dict and existing is not None and existing.credentials:
|
||||
existing_creds = (
|
||||
json.loads(existing.credentials) if isinstance(existing.credentials, str) else dict(existing.credentials)
|
||||
)
|
||||
if _TOKEN_EXCHANGE_COLUMN_FIELDS & existing_creds.keys():
|
||||
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
|
||||
legacy_value = existing_creds.pop(te_field, None)
|
||||
if legacy_value is not None and te_field not in data_dict and getattr(existing, te_field, None) is None:
|
||||
data_dict[te_field] = legacy_value
|
||||
data_dict["credentials"] = safe_dumps(existing_creds)
|
||||
|
||||
# Merge credentials: preserve existing fields not present in the update.
|
||||
# Without this, a partial credential update (e.g. changing only region)
|
||||
# would wipe encrypted secrets that the UI cannot display back.
|
||||
if "credentials" in data_dict and data_dict["credentials"] is not None:
|
||||
if existing and existing.credentials:
|
||||
# Only merge when auth_type is unchanged. Switching auth types
|
||||
# (e.g. oauth2 → api_key) should replace credentials entirely
|
||||
# to avoid stale secrets from the previous auth type lingering.
|
||||
auth_type_unchanged = data.auth_type is None or data.auth_type == existing.auth_type
|
||||
if auth_type_unchanged:
|
||||
# Only merge when the credential CLASS is unchanged. A cross-class switch
|
||||
# (e.g. oauth2 → api_key, or oauth2 → true_passthrough) replaces credentials
|
||||
# entirely to avoid stale secrets from the previous class lingering; a switch
|
||||
# within the client-forwarded class (true_passthrough ↔ oauth_delegate) keeps
|
||||
# the same declared app and so must merge, not replace.
|
||||
if not auth_type_changed:
|
||||
existing_creds = (
|
||||
json.loads(existing.credentials)
|
||||
if isinstance(existing.credentials, str)
|
||||
|
|
@ -636,13 +758,35 @@ async def update_mcp_server(
|
|||
if isinstance(data_dict["credentials"], str)
|
||||
else dict(data_dict["credentials"])
|
||||
)
|
||||
# New values override existing; existing keys not in update are preserved
|
||||
merged = {**existing_creds, **new_creds}
|
||||
# New values override existing; existing keys not in update are preserved. A client
|
||||
# rotation additionally drops the previous app's stale minted token keys.
|
||||
merged = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds)
|
||||
# Migrate-on-write for legacy rows: token-exchange settings the
|
||||
# old blob shape carried move to their dedicated columns (unless
|
||||
# the caller set the column this update, or the row already has
|
||||
# one) and are never re-persisted in the blob. Stored plaintext,
|
||||
# so the merged value lifts as-is.
|
||||
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
|
||||
legacy_value = merged.pop(te_field, None)
|
||||
if (
|
||||
legacy_value is not None
|
||||
and te_field not in data_dict
|
||||
and getattr(existing, te_field, None) is None
|
||||
):
|
||||
data_dict[te_field] = legacy_value
|
||||
data_dict["credentials"] = safe_dumps(merged)
|
||||
|
||||
# Add audit fields
|
||||
data_dict["updated_by"] = touched_by
|
||||
|
||||
# prisma-python rejects a raw ``None`` for a ``Json?`` field ("value is required but not set"); the
|
||||
# clear paths above use ``None`` as the merge-skip sentinel, so translate it here to ``Json(None)``,
|
||||
# which writes SQL null and reads back as ``None``. Done at the edge so the merge guards stay simple.
|
||||
if "credentials" in data_dict and data_dict["credentials"] is None:
|
||||
from prisma import Json # noqa: PLC0415 # local import: prisma may be ungenerated at module load in some tools
|
||||
|
||||
data_dict["credentials"] = Json(None)
|
||||
|
||||
updated_mcp_server = await MCPServerRepository(prisma_client).table.update(
|
||||
where={"server_id": data.server_id},
|
||||
data=data_dict, # type: ignore
|
||||
|
|
@ -998,6 +1142,103 @@ async def list_user_oauth_credentials(
|
|||
return results
|
||||
|
||||
|
||||
def _decrypted_credential_field(creds: Dict[str, object], field: str) -> object:
|
||||
"""Return one credential field decrypted with the global salt key; non-string and legacy
|
||||
plaintext values come back unchanged (decrypt_value_helper returns the original on failure)."""
|
||||
value = creds.get(field)
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
return decrypt_value_helper(
|
||||
value=value,
|
||||
key=field,
|
||||
exception_type="debug",
|
||||
return_original_value=True,
|
||||
)
|
||||
|
||||
|
||||
def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
|
||||
"""The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url, or
|
||||
spec_path for OpenAPI servers), the OAuth mode/grant (auth_type, oauth2_flow), the
|
||||
authorization-server endpoints, and the OAuth client + scopes. Mirrors the dashboard's
|
||||
getOAuthAuthorizationIdentity. When any of these change on a server update, previously stored
|
||||
per-user tokens were minted for the old identity and are stale. Excludes transport and
|
||||
delegate_auth_to_upstream, which do not affect what token is minted (RFC 8707/8693).
|
||||
|
||||
client_id/client_secret are compared decrypted: stored values are NaCl-encrypted with a fresh
|
||||
nonce on every write, so comparing ciphertext would flag every routine save as an identity
|
||||
change and purge tokens that are still valid."""
|
||||
creds = getattr(server, "credentials", None)
|
||||
if isinstance(creds, str):
|
||||
try:
|
||||
parsed: object = json.loads(creds)
|
||||
except ValueError:
|
||||
parsed = None
|
||||
else:
|
||||
parsed = creds
|
||||
creds_dict: Dict[str, object] = parsed if isinstance(parsed, dict) else {}
|
||||
return (
|
||||
getattr(server, "url", None),
|
||||
getattr(server, "spec_path", None),
|
||||
getattr(server, "auth_type", None),
|
||||
getattr(server, "oauth2_flow", None),
|
||||
getattr(server, "authorization_url", None),
|
||||
getattr(server, "token_url", None),
|
||||
getattr(server, "registration_url", None),
|
||||
_decrypted_credential_field(creds_dict, "client_id"),
|
||||
_decrypted_credential_field(creds_dict, "client_secret"),
|
||||
creds_dict.get("scopes"),
|
||||
)
|
||||
|
||||
|
||||
async def purge_user_oauth_credentials_for_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> int:
|
||||
"""Delete every stored per-user OAuth token for a server and invalidate each user's cached
|
||||
token everywhere it can be served from (the legacy per-user token cache and the v2 per-user OAuth
|
||||
token store), so no user keeps a token minted for a superseded configuration. Called when a server
|
||||
update changes a mint-relevant field (see mcp_oauth_token_identity). Returns the number of rows
|
||||
removed.
|
||||
|
||||
LiteLLM_MCPUserCredentials also stores BYOK API keys in the same column; only rows whose payload
|
||||
decodes as an OAuth2 credential (see _decode_oauth_payload) are deleted, because a config change
|
||||
only invalidates minted tokens, never a user's own stored key. Rows are therefore deleted per
|
||||
(user_id, server_id) pair rather than by a blanket server_id filter. An OAuth row inserted while
|
||||
the purge runs for a user not yet enumerated survives; a re-auth completing in the window for an
|
||||
already-enumerated user is deleted along with the stale row (the pair delete cannot tell them
|
||||
apart), which costs that user one extra re-auth and nothing else.
|
||||
|
||||
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
|
||||
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
|
||||
repo = MCPUserCredentialsRepository(prisma_client)
|
||||
rows = await repo.table.find_many(where={"server_id": server_id})
|
||||
oauth_rows = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
|
||||
if not oauth_rows:
|
||||
return 0
|
||||
deleted_count = await repo.table.delete_many(
|
||||
where={"server_id": server_id, "user_id": {"in": [row.user_id for row in oauth_rows]}}
|
||||
)
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
|
||||
for row in oauth_rows:
|
||||
await invalidate_token_cache(row.user_id, server_id)
|
||||
if deleted_count != len(oauth_rows):
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s: purge removed %d OAuth credential row(s) but %d were enumerated; "
|
||||
"row(s) were deleted concurrently during the purge",
|
||||
server_id,
|
||||
deleted_count,
|
||||
len(oauth_rows),
|
||||
)
|
||||
return deleted_count
|
||||
|
||||
|
||||
async def refresh_user_oauth_token(
|
||||
prisma_client: PrismaClient,
|
||||
user_id: str,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
import asyncio
|
||||
import html as _html
|
||||
import json
|
||||
import math
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
|
@ -10,7 +12,8 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|||
import httpx
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from pydantic import BaseModel, SecretStr, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -37,6 +40,10 @@ from litellm.types.mcp import MCPAuth, MCPCredentials
|
|||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
EnvelopeKeys,
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth
|
||||
|
||||
# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers.
|
||||
|
|
@ -326,66 +333,136 @@ def _litellm_key_from_request(request: Request) -> Optional[str]:
|
|||
return None
|
||||
|
||||
|
||||
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]:
|
||||
"""The key's ``user_id``, or ``None`` if the key is blocked or expired.
|
||||
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
|
||||
"""``True`` when the presented key is neither blocked nor past its expiry.
|
||||
|
||||
The OAuth token endpoint is unauthenticated, so the presented key is validated here before its
|
||||
identity is trusted to key a stored credential; a revoked or expired key must not be able to
|
||||
write or overwrite the per-user OAuth token. ``get_key_object`` resolves a row without these
|
||||
checks (the main ``user_api_key_auth`` pipeline enforces them downstream, which this endpoint
|
||||
bypasses), so they are applied here. Deleted keys are already rejected upstream, where
|
||||
``get_key_object`` raises on a row that no longer exists.
|
||||
The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is
|
||||
trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential.
|
||||
``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline
|
||||
enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys
|
||||
are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists.
|
||||
|
||||
This is an active-state gate only; it deliberately does not require a ``user_id``. A valid
|
||||
team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating
|
||||
on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token
|
||||
store) derive it separately via :func:`_active_key_user_id`.
|
||||
|
||||
Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make
|
||||
``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution
|
||||
``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed
|
||||
behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising.
|
||||
"""
|
||||
if key_obj.blocked is True:
|
||||
return None
|
||||
return False
|
||||
expires = key_obj.expires
|
||||
if expires is not None:
|
||||
expiry = expires if isinstance(expires, datetime) else datetime.fromisoformat(expires)
|
||||
if isinstance(expires, datetime):
|
||||
expiry = expires
|
||||
else:
|
||||
try:
|
||||
expiry = datetime.fromisoformat(expires)
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
if expiry < datetime.now(timezone.utc):
|
||||
return None
|
||||
return key_obj.user_id
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _extract_user_id_from_request(request: Request) -> Optional[str]:
|
||||
"""Resolve the LiteLLM ``user_id`` at the OAuth token endpoint so a per-user token is stored
|
||||
under the same identity the egress later reads it by (``user_api_key_auth.user_id``).
|
||||
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None:
|
||||
"""The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no
|
||||
``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which
|
||||
needs a user to key the stored credential; the bridge mint uses the key hash and does not."""
|
||||
return key_obj.user_id if _key_is_active(key_obj) else None
|
||||
|
||||
Resolves authoritatively via ``get_key_object`` (cache first, then DB) instead of a raw cache
|
||||
peek. On a multi-replica gateway the token-exchange request can land on a worker whose in-memory
|
||||
cache never saw the key, and a cross-replica Redis hit deserializes to a plain ``dict`` rather
|
||||
than a ``UserAPIKeyAuth``; the previous code read only ``Authorization`` and did
|
||||
``getattr(cached, "user_id")`` with no ``model_type`` rehydration and no DB fallback, so it
|
||||
silently returned ``None`` and the token was never persisted, which makes the egress 401 on every
|
||||
reconnect. The resolved key is validated (``_active_key_user_id``) before its identity is trusted,
|
||||
so a blocked or expired key cannot write. Returns ``None`` when no key is present, the key cannot
|
||||
be resolved, or it is blocked/expired.
|
||||
"""
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ResolvedKey:
|
||||
"""An active litellm key resolved from the token request: its hash (the value ``get_key_object``
|
||||
and the cache/DB layer key the record by) and the live record."""
|
||||
|
||||
key_hash: str
|
||||
key: "UserAPIKeyAuth"
|
||||
|
||||
|
||||
_KeyResolutionFailure = Literal["no_active_key", "unavailable", "unresolvable"]
|
||||
"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully
|
||||
instead of blaming the client for a gateway problem:
|
||||
- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the
|
||||
caller's request is at fault)
|
||||
- ``unavailable``: the auth database was transiently unreachable while resolving (retryable)
|
||||
- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected
|
||||
error) -- a gateway fault, not the caller's
|
||||
The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission
|
||||
(egress) never disagree on the status of the same outage."""
|
||||
|
||||
|
||||
async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure":
|
||||
"""Resolve the presented litellm key to an active key record, or say precisely why not.
|
||||
|
||||
Single resolution path the OAuth token endpoint reuses, resolving authoritatively via
|
||||
``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller
|
||||
can tell "the client sent no usable credential" (a request error) apart from "the gateway could not
|
||||
check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let
|
||||
a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or
|
||||
expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``)
|
||||
resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway
|
||||
fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key,
|
||||
a database-service-unavailable error is a retryable outage, and anything else is an unexpected
|
||||
gateway fault."""
|
||||
token = _litellm_key_from_request(request)
|
||||
if not token:
|
||||
return None
|
||||
try:
|
||||
from litellm.proxy._types import hash_token # noqa: PLC0415
|
||||
from litellm.proxy.auth.auth_checks import get_key_object # noqa: PLC0415
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
return "no_active_key"
|
||||
from litellm.proxy._types import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
ProxyException,
|
||||
hash_token,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
get_key_object,
|
||||
)
|
||||
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
PrismaDBExceptionHandler,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return "unresolvable"
|
||||
key_hash = hash_token(token)
|
||||
try:
|
||||
key_obj = await get_key_object(
|
||||
hashed_token=hash_token(token),
|
||||
hashed_token=key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
return _active_key_user_id(key_obj)
|
||||
except Exception as exc:
|
||||
except (ProxyException, HTTPException):
|
||||
return "no_active_key"
|
||||
except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc):
|
||||
return "unavailable"
|
||||
verbose_logger.debug(
|
||||
"_extract_user_id_from_request: could not resolve a LiteLLM user_id for the presented "
|
||||
"key (%s); per-user token will not be stored server-side.",
|
||||
"_resolve_active_litellm_key: unexpected key-resolution error (%s)",
|
||||
type(exc).__name__,
|
||||
)
|
||||
return "unresolvable"
|
||||
if not _key_is_active(key_obj):
|
||||
return "no_active_key"
|
||||
return _ResolvedKey(key_hash=key_hash, key=key_obj)
|
||||
|
||||
|
||||
async def _extract_user_id_from_request(request: Request) -> str | None:
|
||||
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
|
||||
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
|
||||
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
|
||||
the bridge mint, which must status those outcomes differently, consumes
|
||||
:func:`_resolve_active_litellm_key` directly."""
|
||||
resolved = await _resolve_active_litellm_key(request)
|
||||
if not isinstance(resolved, _ResolvedKey):
|
||||
return None
|
||||
return _active_key_user_id(resolved.key)
|
||||
|
||||
|
||||
async def _store_per_user_token_server_side(
|
||||
|
|
@ -448,6 +525,12 @@ async def _store_per_user_token_server_side(
|
|||
)
|
||||
return # Don't warm Redis if DB write failed
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id)
|
||||
|
||||
# Warm the Redis cache so the first subsequent MCP call is a cache hit
|
||||
ttl = _compute_per_user_token_ttl(server, expires_in)
|
||||
await mcp_per_user_token_cache.set(
|
||||
|
|
@ -459,8 +542,20 @@ async def _store_per_user_token_server_side(
|
|||
|
||||
|
||||
def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
|
||||
"""Reject a non-oauth2 server from the gateway's OAuth authorize/token/register flow."""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2:
|
||||
"""Reject a server without upstream OAuth from the gateway's authorize/token/register flow.
|
||||
|
||||
The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are allowed
|
||||
through: the caller owns the upstream token, and this relayed flow is how a browser obtains
|
||||
one against the upstream IdP (the admin UI's browser-only Authorize uses it). The minted
|
||||
token is upstream-audienced and held by the caller; the gateway persists nothing for these
|
||||
modes (``_persist_dcr_client_registration`` skips them unconditionally, so even the admin
|
||||
Authorize path with ``persist_credentials`` enabled writes nothing to the server row).
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
)
|
||||
|
||||
if mcp_server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -481,23 +576,86 @@ def _raise_unless_oauth2_discovery_server(
|
|||
mcp_server_name: Optional[str],
|
||||
description: str,
|
||||
) -> None:
|
||||
"""404 a NAMED discovery request unless it resolves to an oauth2 server.
|
||||
"""404 a NAMED discovery request unless it resolves to an oauth2 or DCR-bridge server.
|
||||
|
||||
A named server that is unknown (or hidden from the caller) and one that exists
|
||||
but is non-oauth2 both return the same 404, so the well-known discovery paths
|
||||
cannot be used to enumerate non-OAuth server names. Root discovery (no name) is
|
||||
unaffected, and pass-through servers are resolved by the caller before this runs.
|
||||
DCR-bridge servers are admitted because they serve the gateway's own authorization
|
||||
server metadata (the register, authorize, and token relays).
|
||||
"""
|
||||
if mcp_server_name is None:
|
||||
return
|
||||
if mcp_server is not None and mcp_server.auth_type == MCPAuth.oauth2:
|
||||
return
|
||||
if mcp_server is not None and mcp_server.is_dcr_bridge:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"MCP server '{mcp_server_name}' is {description}",
|
||||
)
|
||||
|
||||
|
||||
def _dcr_bridge_relays_client_registration(mcp_server: MCPServer) -> bool:
|
||||
"""True when a DCR-bridge server relays client registration to the upstream authorization
|
||||
server instead of short-circuiting to an admin-configured OAuth client. In the relay arm the
|
||||
upstream holds each client's own registration, so the authorize and token relays pass the
|
||||
client's ``client_id`` and ``redirect_uri`` through verbatim and the authorization code
|
||||
returns directly to the client's redirect URI without transiting the gateway. Gateway-side
|
||||
redirect trust and the ``/callback`` state relay therefore only apply to the short-circuit
|
||||
arm, where the upstream only knows the gateway's own callback."""
|
||||
return mcp_server.is_dcr_bridge and bool(mcp_server.registration_url) and not mcp_server.client_id
|
||||
|
||||
|
||||
def _require_s256_pkce(
|
||||
code_challenge: Optional[str],
|
||||
code_challenge_method: Optional[str],
|
||||
) -> Tuple[str, str]:
|
||||
"""DCR-bridge servers serve unauthenticated public OAuth clients, so the PKCE downgrade
|
||||
paths (no challenge, or a non-S256 method; RFC 7636 defaults a missing method to ``plain``)
|
||||
are rejected at the gateway instead of relying on upstream enforcement. Returns the
|
||||
validated pair so callers get non-optional values."""
|
||||
if code_challenge and code_challenge_method == "S256":
|
||||
return code_challenge, code_challenge_method
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"This server requires PKCE: send code_challenge with "
|
||||
"code_challenge_method=S256 on the authorization request"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _redirect_to_upstream_authorize(
|
||||
*,
|
||||
mcp_server: MCPServer,
|
||||
client_id: str,
|
||||
redirect_uri: str,
|
||||
state: str,
|
||||
code_challenge: str,
|
||||
code_challenge_method: str,
|
||||
response_type: Optional[str],
|
||||
scope: Optional[str],
|
||||
) -> RedirectResponse:
|
||||
"""The bridge relay arm's authorize redirect: every client-supplied parameter passes through
|
||||
to the upstream authorize endpoint verbatim, no relay state cookie is set, and the upstream
|
||||
enforces its own registered redirect binding for the client."""
|
||||
scope_value = scope or (" ".join(mcp_server.scopes) if mcp_server.scopes else None)
|
||||
passthrough_params = {
|
||||
"client_id": client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
"state": state,
|
||||
"response_type": response_type or "code",
|
||||
"code_challenge": code_challenge,
|
||||
"code_challenge_method": code_challenge_method,
|
||||
**({"scope": scope_value} if scope_value else {}),
|
||||
}
|
||||
parsed_auth_url = urlparse(mcp_server.authorization_url or "")
|
||||
merged_params = {**dict(parse_qsl(parsed_auth_url.query)), **passthrough_params}
|
||||
return RedirectResponse(urlunparse(parsed_auth_url._replace(query=urlencode(merged_params))))
|
||||
|
||||
|
||||
async def authorize_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -509,11 +667,28 @@ async def authorize_with_server(
|
|||
response_type: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
):
|
||||
if mcp_server.auth_type != "oauth2":
|
||||
raise HTTPException(status_code=400, detail="MCP server is not OAuth2")
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
if mcp_server.authorization_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server authorization url is not set")
|
||||
|
||||
if mcp_server.is_dcr_bridge:
|
||||
# Enforce S256 PKCE on both bridge arms. The relay arm forwards the validated,
|
||||
# now-non-optional pair to the upstream authorize; the short-circuit arm keeps
|
||||
# calling this for its enforcement side effect, then falls through to the gateway
|
||||
# /callback flow below, which reads the original code_challenge names.
|
||||
bridge_challenge, bridge_method = _require_s256_pkce(code_challenge, code_challenge_method)
|
||||
if _dcr_bridge_relays_client_registration(mcp_server):
|
||||
return _redirect_to_upstream_authorize(
|
||||
mcp_server=mcp_server,
|
||||
client_id=client_id,
|
||||
redirect_uri=redirect_uri,
|
||||
state=state,
|
||||
code_challenge=bridge_challenge,
|
||||
code_challenge_method=bridge_method,
|
||||
response_type=response_type,
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Trusted redirect_uri: same-origin, loopback, or ops-allowlisted.
|
||||
# The URI is encrypted into the OAuth state and decoded on
|
||||
# /callback to redirect the user back; a non-trusted URI would be
|
||||
|
|
@ -556,6 +731,255 @@ async def authorize_with_server(
|
|||
return response
|
||||
|
||||
|
||||
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
|
||||
"""Why an upstream token response cannot back a bridge envelope:
|
||||
- ``no_access_token``: the response carries no usable ``access_token``
|
||||
- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream
|
||||
token that is already dead, so sealing it would forward a bearer the edge cannot use
|
||||
An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the
|
||||
envelope caps it, the by-design behaviour for an upstream that omits the field."""
|
||||
|
||||
|
||||
def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']":
|
||||
"""Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent
|
||||
or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports
|
||||
as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is
|
||||
already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h
|
||||
cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
|
||||
positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the
|
||||
envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded
|
||||
(an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` /
|
||||
``Infinity`` / oversized input, which reads as unparseable rather than surfacing as a 500."""
|
||||
if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)):
|
||||
return "unspecified"
|
||||
try:
|
||||
numeric = float(raw_expires_in)
|
||||
seconds = int(numeric)
|
||||
except (ValueError, TypeError, OverflowError):
|
||||
return "unspecified"
|
||||
if numeric <= 0:
|
||||
return "expired"
|
||||
return max(1, seconds)
|
||||
|
||||
|
||||
def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection":
|
||||
"""Validate an upstream OAuth token response into a typed grant, or say why it cannot back an
|
||||
envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the
|
||||
grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown
|
||||
lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is
|
||||
honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to
|
||||
the cap."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
|
||||
if not isinstance(token_response, dict):
|
||||
return "no_access_token"
|
||||
access = token_response.get("access_token")
|
||||
if not isinstance(access, str) or not access:
|
||||
return "no_access_token"
|
||||
lifetime = _classify_upstream_lifetime(token_response.get("expires_in"))
|
||||
if lifetime == "expired":
|
||||
return "expired_lifetime"
|
||||
token_type = token_response.get("token_type")
|
||||
scope = token_response.get("scope")
|
||||
return UpstreamTokenGrant(
|
||||
access_token=SecretStr(access),
|
||||
token_type=token_type if isinstance(token_type, str) and token_type else "Bearer",
|
||||
# The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards
|
||||
# only token_type + access_token), so it would be dead weight embedding a long-lived upstream
|
||||
# credential in the client-held bearer, and it enlarges the envelope. Refresh support is a
|
||||
# follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap.
|
||||
refresh_token=None,
|
||||
scope=scope if isinstance(scope, str) and scope else None,
|
||||
expires_in=lifetime if isinstance(lifetime, int) else None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DCR-bridge oauth_delegate mint: a three-phase pipeline whose failures are values.
|
||||
#
|
||||
# prepare (before the upstream exchange) -> validate every precondition and resolve identity+keys
|
||||
# exchange (the single-use upstream code is consumed here, in exchange_token_with_server)
|
||||
# finish (after the exchange) -> seal the upstream grant into the client-held envelope
|
||||
#
|
||||
# Every precondition lives in ``prepare``, which runs BEFORE the exchange, so no failure can burn the
|
||||
# single-use code or rotate a refresh token, for either grant type -- that whole class of bug is gone
|
||||
# by construction rather than guarded case by case. Failures are values mapped to an OAuth-shaped
|
||||
# response in one place (``_bridge_mint_error_response``), so status codes and the RFC 6749 §5.2 body
|
||||
# shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_BridgeMintError = Literal[
|
||||
"no_identity",
|
||||
"unsupported_grant",
|
||||
"identity_unavailable",
|
||||
"identity_unresolvable",
|
||||
"not_configured",
|
||||
"no_upstream_token",
|
||||
"upstream_token_expired",
|
||||
"too_large",
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BridgeMintReady:
|
||||
"""Everything the seal needs, resolved once before the exchange: the authorizing key hash and the
|
||||
master-key-derived envelope keys. Passing this forward means identity resolution and key derivation
|
||||
happen exactly once, and ``_finish_bridge_mint`` has no preconditions left that could fail."""
|
||||
|
||||
key_hash: str
|
||||
keys: "EnvelopeKeys"
|
||||
|
||||
|
||||
def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
|
||||
"""Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape
|
||||
(top-level ``error``, no-store headers) for every case, with a status truthful about where the
|
||||
failure is. The caller's request is 400, a transient gateway outage is 503, a gateway
|
||||
misconfiguration is 500, and an upstream problem is 502. The identity-resolution statuses match how
|
||||
admission statuses the same conditions on the egress side, so mint and admit never disagree under
|
||||
one outage."""
|
||||
match error:
|
||||
case "no_identity":
|
||||
status, code, desc = (
|
||||
400,
|
||||
"invalid_request",
|
||||
"this server issues a gateway-bound credential; send a litellm credential "
|
||||
"(x-litellm-api-key or Authorization) on the token request",
|
||||
)
|
||||
case "unsupported_grant":
|
||||
status, code, desc = (
|
||||
400,
|
||||
"unsupported_grant_type",
|
||||
"this server issues a gateway-bound credential and supports only the authorization_code "
|
||||
"grant; re-run authorization_code to renew rather than refresh_token",
|
||||
)
|
||||
case "identity_unavailable":
|
||||
status, code, desc = (
|
||||
503,
|
||||
"temporarily_unavailable",
|
||||
"the authentication database is temporarily unreachable; retry shortly",
|
||||
)
|
||||
case "identity_unresolvable":
|
||||
status, code, desc = (
|
||||
500,
|
||||
"server_error",
|
||||
"the gateway could not resolve the litellm identity for this request",
|
||||
)
|
||||
case "not_configured":
|
||||
status, code, desc = (
|
||||
500,
|
||||
"server_error",
|
||||
"the gateway is not configured to mint a gateway-bound credential (master_key is not set)",
|
||||
)
|
||||
case "no_upstream_token":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token response has no usable access_token",
|
||||
)
|
||||
case "upstream_token_expired":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token response reports an already-expired lifetime",
|
||||
)
|
||||
case "too_large":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token is too large to seal into a gateway-bound credential",
|
||||
)
|
||||
case _:
|
||||
assert_never(error)
|
||||
return JSONResponse(
|
||||
status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS
|
||||
)
|
||||
|
||||
|
||||
def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
|
||||
"""Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays
|
||||
truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that
|
||||
cannot resolve identity is 500."""
|
||||
match failure:
|
||||
case "no_active_key":
|
||||
return "no_identity"
|
||||
case "unavailable":
|
||||
return "identity_unavailable"
|
||||
case "unresolvable":
|
||||
return "identity_unresolvable"
|
||||
case _:
|
||||
assert_never(failure)
|
||||
|
||||
|
||||
def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _BridgeMintError:
|
||||
"""Lift an upstream-response rejection into the mint taxonomy; both are upstream faults (502)."""
|
||||
match rejection:
|
||||
case "no_access_token":
|
||||
return "no_upstream_token"
|
||||
case "expired_lifetime":
|
||||
return "upstream_token_expired"
|
||||
case _:
|
||||
assert_never(rejection)
|
||||
|
||||
|
||||
async def _prepare_bridge_mint(request: Request, grant_type: str) -> "_BridgeMintReady | _BridgeMintError":
|
||||
"""Phase 1, BEFORE the upstream exchange: reject a grant this mint does not support, confirm the
|
||||
gateway can mint (master_key set), resolve the litellm identity, and derive the envelope keys.
|
||||
Returns a ready context or a precise failure value. Running before the exchange is what makes every
|
||||
failure here fail closed without consuming the single-use code or rotating a refresh token. A bridge
|
||||
server issues only envelopes and seals no upstream refresh_token, so the client holds none to
|
||||
present: the refresh_token grant is rejected up front rather than exchanged (which could rotate the
|
||||
upstream credential) and its result then discarded. Identity-resolution failures keep their origin
|
||||
so the mapper statuses each truthfully."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
envelope_keys_from_master_key,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
master_key,
|
||||
)
|
||||
|
||||
if grant_type != "authorization_code":
|
||||
return "unsupported_grant"
|
||||
if not master_key:
|
||||
return "not_configured"
|
||||
resolved = await _resolve_active_litellm_key(request)
|
||||
if not isinstance(resolved, _ResolvedKey):
|
||||
return _key_resolution_failure_to_mint_error(resolved)
|
||||
return _BridgeMintReady(key_hash=resolved.key_hash, keys=envelope_keys_from_master_key(master_key))
|
||||
|
||||
|
||||
def _finish_bridge_mint(
|
||||
ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime
|
||||
) -> "JSONResponse | _BridgeMintError":
|
||||
"""Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held envelope using
|
||||
the pre-resolved identity and keys, so the client holds one bearer that admits it and forwards the
|
||||
upstream token with nothing stored server-side. The only failures here are properties of the
|
||||
upstream response (no usable token, an already-expired lifetime, or a token too large to seal),
|
||||
returned as values."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
build_bridge_token_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
EnvelopeIdentity,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
|
||||
grant = _bridge_grant_from_token_response(token_response)
|
||||
if not isinstance(grant, UpstreamTokenGrant):
|
||||
return _upstream_rejection_to_mint_error(grant)
|
||||
identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=ready.key_hash)
|
||||
sealed = build_bridge_token_response(identity, grant, ready.keys, now)
|
||||
if not isinstance(sealed, SealedEnvelope):
|
||||
return "too_large"
|
||||
# Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the
|
||||
# client is never told the bearer lives past the point admission (which uses that exp) rejects it.
|
||||
expires_in = max(0, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp()))
|
||||
body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in}
|
||||
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
|
||||
|
||||
async def exchange_token_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -575,8 +999,12 @@ async def exchange_token_with_server(
|
|||
if mcp_server.token_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server token url is not set")
|
||||
|
||||
# The id and secret must come from the same source. When the server-side client_id wins,
|
||||
# falling back to the caller's secret pairs the persisted client with a foreign secret; the
|
||||
# register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a
|
||||
# persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s.
|
||||
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
|
||||
resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret
|
||||
resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=mcp_server.token_endpoint_auth_method,
|
||||
|
|
@ -605,16 +1033,36 @@ async def exchange_token_with_server(
|
|||
status_code=400,
|
||||
detail="code is required for authorization_code grant",
|
||||
)
|
||||
bridge_token_relay = _dcr_bridge_relays_client_registration(mcp_server)
|
||||
if bridge_token_relay and not redirect_uri:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"redirect_uri is required for the authorization_code grant on this server; "
|
||||
"send the same redirect_uri used on the authorization request"
|
||||
),
|
||||
)
|
||||
proxy_base_url = get_request_base_url(request)
|
||||
resolved_redirect_uri = redirect_uri if bridge_token_relay else f"{proxy_base_url}/callback"
|
||||
token_data = {
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": f"{proxy_base_url}/callback",
|
||||
"redirect_uri": resolved_redirect_uri,
|
||||
**client_auth.body,
|
||||
}
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
# Phase 1: for a bridge oauth_delegate mint, validate all preconditions and resolve identity+keys
|
||||
# BEFORE the exchange below consumes the single-use upstream code, and carry the ready context to
|
||||
# phase 3. A failure here returns without ever touching the upstream credential.
|
||||
bridge_mint_ready: _BridgeMintReady | None = None
|
||||
if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge:
|
||||
prepared = await _prepare_bridge_mint(request, grant_type)
|
||||
if not isinstance(prepared, _BridgeMintReady):
|
||||
return _bridge_mint_error_response(prepared)
|
||||
bridge_mint_ready = prepared
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
|
|
@ -627,9 +1075,18 @@ async def exchange_token_with_server(
|
|||
detail="MCP upstream token endpoint returned no response",
|
||||
)
|
||||
|
||||
response.raise_for_status()
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
if "invalid_target" in exc.response.text:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: the upstream authorization server rejected the token request with "
|
||||
"invalid_target; it may require RFC 8707 resource indicators, which the gateway "
|
||||
"does not send yet (tracked as LIT-4339)",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
raise
|
||||
token_response = response.json()
|
||||
access_token = token_response["access_token"]
|
||||
|
||||
# Validate token response against server-configured rules before any storage.
|
||||
# This rejects tokens from wrong Slack workspaces, Atlassian orgs, etc.
|
||||
|
|
@ -669,8 +1126,17 @@ async def exchange_token_with_server(
|
|||
mcp_server.server_id,
|
||||
)
|
||||
|
||||
# A DCR-bridge oauth_delegate server hands the client a gateway-bound envelope (identity plus the
|
||||
# upstream token) instead of the raw upstream token, so the one bearer both admits the caller and
|
||||
# forwards the upstream credential. Only this mode mints; every other server returns the raw token.
|
||||
if bridge_mint_ready is not None:
|
||||
# Phase 3: seal the upstream grant into the client-held envelope; failures map through the same
|
||||
# OAuth-shaped response as the phase-1 preconditions.
|
||||
minted = _finish_bridge_mint(bridge_mint_ready, mcp_server, token_response, datetime.now(timezone.utc))
|
||||
return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted)
|
||||
|
||||
result = {
|
||||
"access_token": access_token,
|
||||
"access_token": token_response["access_token"],
|
||||
"token_type": token_response.get("token_type", "Bearer"),
|
||||
}
|
||||
|
||||
|
|
@ -698,6 +1164,22 @@ class _PersistedDcrCredentials(BaseModel):
|
|||
client_id: Optional[str] = None
|
||||
client_secret: Optional[str] = None
|
||||
token_endpoint_auth_method: Optional[str] = None
|
||||
redirect_uris: Optional[list[str]] = None
|
||||
|
||||
|
||||
def _redirect_uri_not_registered(credentials: _PersistedDcrCredentials, current_redirect_uri: str) -> bool:
|
||||
"""Whether a persisted DCR client is positively known NOT to cover the current callback.
|
||||
|
||||
A DCR client is bound to the redirect_uris it was registered with; if the proxy's
|
||||
resolved public origin has since changed, every authorize built for it will be
|
||||
rejected by the IdP. Clients persisted before ``redirect_uris`` was recorded (and
|
||||
admin-configured clients, which never get a recording) return False so they are
|
||||
grandfathered rather than re-registered, because re-minting a client_id orphans
|
||||
every user's refresh tokens for that server."""
|
||||
recorded = credentials.redirect_uris
|
||||
if not recorded:
|
||||
return False
|
||||
return current_redirect_uri not in recorded
|
||||
|
||||
|
||||
def _get_persisted_dcr_credentials(credentials: object) -> Optional[_PersistedDcrCredentials]:
|
||||
|
|
@ -764,11 +1246,23 @@ async def _get_persisted_mcp_server_with_dcr_client_id(
|
|||
return persisted_mcp_server, credentials
|
||||
|
||||
|
||||
async def _reuse_persisted_dcr_client_if_available(mcp_server: MCPServer) -> bool:
|
||||
async def _reuse_persisted_dcr_client_if_available(
|
||||
mcp_server: MCPServer, current_redirect_uri: Optional[str] = None
|
||||
) -> bool:
|
||||
persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
|
||||
if persisted is None:
|
||||
return False
|
||||
persisted_mcp_server, credentials = persisted
|
||||
if current_redirect_uri is not None and _redirect_uri_not_registered(credentials, current_redirect_uri):
|
||||
verbose_logger.debug(
|
||||
"register_client_with_server: not reusing persisted DCR client for server_id=%s; its registered "
|
||||
"redirect_uris=%s do not include the current callback %s. The operator-facing warning for this "
|
||||
"re-registration event is emitted once by _persisted_dcr_redirect_uri_is_stale.",
|
||||
mcp_server.server_id,
|
||||
credentials.redirect_uris,
|
||||
current_redirect_uri,
|
||||
)
|
||||
return False
|
||||
if not _apply_persisted_dcr_credentials(mcp_server, credentials):
|
||||
return False
|
||||
|
||||
|
|
@ -787,11 +1281,36 @@ async def _reuse_persisted_dcr_client_if_available(mcp_server: MCPServer) -> boo
|
|||
return bool(mcp_server.client_id)
|
||||
|
||||
|
||||
DcrRegistrationPersistenceResult = Literal["persisted", "reused", "failed"]
|
||||
async def _persisted_dcr_redirect_uri_is_stale(mcp_server: MCPServer, current_redirect_uri: str) -> bool:
|
||||
"""Whether the server's persisted DCR client is bound to redirect_uris that no longer
|
||||
cover the current proxy callback, meaning authorize is guaranteed to fail IdP-side.
|
||||
|
||||
Consulted when the in-memory server already carries a hydrated client_id, which
|
||||
otherwise short-circuits registration before any redirect check can run. Servers
|
||||
without a persisted DCR recording (admin-configured client_id, or registered before
|
||||
redirect_uris were recorded) are never reported stale."""
|
||||
persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
|
||||
if persisted is None:
|
||||
return False
|
||||
_, credentials = persisted
|
||||
if not _redirect_uri_not_registered(credentials, current_redirect_uri):
|
||||
return False
|
||||
verbose_logger.warning(
|
||||
"register_client_with_server: persisted DCR client for server_id=%s is registered with redirect_uris=%s "
|
||||
"which do not include the current callback %s (proxy origin changed); registering a replacement client. "
|
||||
"Users previously signed in to this server will need to re-authenticate.",
|
||||
mcp_server.server_id,
|
||||
credentials.redirect_uris,
|
||||
current_redirect_uri,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
DcrRegistrationPersistenceResult = Literal["persisted", "reused", "skipped", "failed"]
|
||||
|
||||
|
||||
async def _persist_dcr_client_registration(
|
||||
mcp_server: MCPServer, registration_response: object
|
||||
mcp_server: MCPServer, registration_response: object, current_redirect_uri: str
|
||||
) -> DcrRegistrationPersistenceResult:
|
||||
"""Persist the dynamically registered OAuth client (RFC 7591) onto the MCP server row.
|
||||
|
||||
|
|
@ -801,7 +1320,23 @@ async def _persist_dcr_client_registration(
|
|||
full re-authorization instead of a silent refresh. Mirrors the ``encrypt_credentials``
|
||||
write that ``client_credentials`` and token exchange already use. Failures are logged,
|
||||
never raised: registration still returns to the caller even when persistence fails.
|
||||
|
||||
The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are skipped
|
||||
unconditionally: the caller holds the upstream token and the gateway must hold no OAuth
|
||||
client identity for these servers. Persisting here would stamp ``oauth2_flow`` and a
|
||||
``client_id`` onto a server whose mode promises the gateway stores nothing, making a
|
||||
fresh pass-through server read as gateway-authorized.
|
||||
|
||||
``redirect_uris`` records what the client is bound to so a later origin change can be
|
||||
detected as a positive mismatch and trigger re-registration instead of stranding the
|
||||
server on IdP-side redirect_uri rejections. ``client_secret`` and
|
||||
``token_endpoint_auth_method`` are written explicitly (None when absent) because
|
||||
``update_mcp_server`` merges credential blobs: a re-registered public client must not
|
||||
inherit the previous client's secret or auth method.
|
||||
"""
|
||||
if mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate:
|
||||
return "skipped"
|
||||
|
||||
try:
|
||||
registration = _DcrClientRegistration.model_validate(registration_response)
|
||||
except ValidationError as exc:
|
||||
|
|
@ -813,17 +1348,16 @@ async def _persist_dcr_client_registration(
|
|||
)
|
||||
return "failed"
|
||||
|
||||
if await _reuse_persisted_dcr_client_if_available(mcp_server):
|
||||
if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri=current_redirect_uri):
|
||||
return "reused"
|
||||
|
||||
credentials: MCPCredentials = {
|
||||
"client_id": registration.client_id,
|
||||
**({"client_secret": registration.client_secret} if registration.client_secret is not None else {}),
|
||||
**(
|
||||
{"token_endpoint_auth_method": "client_secret_basic"}
|
||||
if registration.token_endpoint_auth_method == "client_secret_basic"
|
||||
else {}
|
||||
"client_secret": registration.client_secret,
|
||||
"token_endpoint_auth_method": (
|
||||
"client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None
|
||||
),
|
||||
"redirect_uris": [current_redirect_uri],
|
||||
}
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.db import update_mcp_server # noqa: PLC0415
|
||||
|
|
@ -858,6 +1392,21 @@ async def _persist_dcr_client_registration(
|
|||
return "failed"
|
||||
|
||||
|
||||
_MAX_UPSTREAM_ERROR_CHARS = 500
|
||||
|
||||
|
||||
def _safe_upstream_error_detail(response: httpx.Response) -> str:
|
||||
"""Bounded plaintext summary of an upstream registration failure for the client.
|
||||
|
||||
RFC 7591 error bodies are small JSON objects (``error`` / ``error_description``); relaying the
|
||||
text lets the client read the real reason instead of a bare 500, and the length bound keeps a
|
||||
hostile or oversized upstream body from bloating the gateway response."""
|
||||
body = response.text
|
||||
if not body:
|
||||
return response.reason_phrase or "upstream registration failed"
|
||||
return body[:_MAX_UPSTREAM_ERROR_CHARS]
|
||||
|
||||
|
||||
async def register_client_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -867,19 +1416,28 @@ async def register_client_with_server(
|
|||
token_endpoint_auth_method: Optional[str],
|
||||
fallback_client_id: Optional[str] = None,
|
||||
persist_credentials: bool = False,
|
||||
client_redirect_uris: Optional[list] = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
request_base_url = get_request_base_url(request)
|
||||
current_redirect_uri = f"{request_base_url}/callback"
|
||||
dummy_return = {
|
||||
"client_id": fallback_client_id or mcp_server.server_name,
|
||||
"client_secret": "dummy",
|
||||
"redirect_uris": [f"{request_base_url}/callback"],
|
||||
"redirect_uris": [current_redirect_uri],
|
||||
}
|
||||
|
||||
if mcp_server.client_id:
|
||||
if mcp_server.client_id and not (
|
||||
persist_credentials
|
||||
and mcp_server.registration_url
|
||||
and await _persisted_dcr_redirect_uri_is_stale(mcp_server, current_redirect_uri)
|
||||
):
|
||||
return dummy_return
|
||||
|
||||
if await _reuse_persisted_dcr_client_if_available(mcp_server):
|
||||
if await _reuse_persisted_dcr_client_if_available(
|
||||
mcp_server,
|
||||
current_redirect_uri=current_redirect_uri if persist_credentials else None,
|
||||
):
|
||||
return dummy_return
|
||||
|
||||
if mcp_server.authorization_url is None:
|
||||
|
|
@ -888,12 +1446,19 @@ async def register_client_with_server(
|
|||
if mcp_server.registration_url is None:
|
||||
return dummy_return
|
||||
|
||||
bridge_relay = _dcr_bridge_relays_client_registration(mcp_server)
|
||||
if bridge_relay and not client_redirect_uris:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="redirect_uris is required to register a client with this server",
|
||||
)
|
||||
|
||||
register_data = {
|
||||
"client_name": client_name,
|
||||
"redirect_uris": [f"{request_base_url}/callback"],
|
||||
"grant_types": grant_types or [],
|
||||
"response_types": response_types or [],
|
||||
"token_endpoint_auth_method": token_endpoint_auth_method or "",
|
||||
"redirect_uris": client_redirect_uris if bridge_relay else [current_redirect_uri],
|
||||
"grant_types": grant_types or (["authorization_code", "refresh_token"] if bridge_relay else []),
|
||||
"response_types": response_types or (["code"] if bridge_relay else []),
|
||||
"token_endpoint_auth_method": token_endpoint_auth_method or ("none" if bridge_relay else ""),
|
||||
}
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -911,12 +1476,14 @@ async def register_client_with_server(
|
|||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no response",
|
||||
)
|
||||
if bridge_relay and response.status_code >= 400:
|
||||
raise HTTPException(status_code=response.status_code, detail=_safe_upstream_error_detail(response))
|
||||
response.raise_for_status()
|
||||
|
||||
token_response = response.json()
|
||||
|
||||
if persist_credentials:
|
||||
persistence_result = await _persist_dcr_client_registration(mcp_server, token_response)
|
||||
if persist_credentials and not bridge_relay:
|
||||
persistence_result = await _persist_dcr_client_registration(mcp_server, token_response, current_redirect_uri)
|
||||
if persistence_result == "reused":
|
||||
return dummy_return
|
||||
|
||||
|
|
@ -1292,11 +1859,15 @@ async def _build_oauth_protected_resource_response(
|
|||
"""
|
||||
Build OAuth protected resource response with the appropriate URL pattern.
|
||||
|
||||
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``), the
|
||||
gateway proxies the upstream's own ``oauth-protected-resource`` metadata
|
||||
so that standards-compliant MCP clients discover the **upstream** IdP
|
||||
instead of the gateway. The ``resource`` field is rewritten to the
|
||||
gateway's own URL so clients present the bearer token back to the gateway.
|
||||
For pass-through MCP servers, the gateway proxies the upstream's own
|
||||
``oauth-protected-resource`` metadata so standards-compliant MCP clients
|
||||
discover the **upstream** IdP instead of the gateway. For ``true_passthrough``
|
||||
and ``oauth_delegate`` the metadata is returned verbatim (``resource`` stays
|
||||
the upstream): the caller's token is forwarded to and validated by the
|
||||
upstream, so its audience must be the upstream — rewriting it to the gateway
|
||||
would make a strict IdP (e.g. Entra) refuse to mint it or the upstream reject
|
||||
it. Only the legacy ``is_oauth_passthrough`` opt-in rewrites ``resource`` to
|
||||
the gateway's own URL so clients present the bearer token back to the gateway.
|
||||
|
||||
Args:
|
||||
request: FastAPI Request object
|
||||
|
|
@ -1335,9 +1906,18 @@ async def _build_oauth_protected_resource_response(
|
|||
else:
|
||||
resource_url = f"{request_base_url}/mcp"
|
||||
|
||||
if mcp_server is not None and mcp_server_name and mcp_server.is_dcr_bridge:
|
||||
return {
|
||||
"authorization_servers": [f"{request_base_url}/{mcp_server_name}"],
|
||||
"resource": resource_url,
|
||||
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
|
||||
}
|
||||
|
||||
# Pass-through branch: proxy the upstream's own metadata so discovery
|
||||
# directs the client at the real IdP (Okta, Keycloak, …) instead of us.
|
||||
if mcp_server is not None and mcp_server.is_oauth_passthrough:
|
||||
if mcp_server is not None and (
|
||||
mcp_server.is_oauth_passthrough or mcp_server.is_oauth_delegate or mcp_server.is_true_passthrough
|
||||
):
|
||||
try:
|
||||
upstream_metadata = await fetch_upstream_oauth_protected_resource(mcp_server)
|
||||
except Exception as exc:
|
||||
|
|
@ -1353,8 +1933,9 @@ async def _build_oauth_protected_resource_response(
|
|||
)
|
||||
|
||||
if upstream_metadata is not None:
|
||||
response = {**upstream_metadata, "resource": resource_url}
|
||||
return response
|
||||
if mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate:
|
||||
return upstream_metadata
|
||||
return {**upstream_metadata, "resource": resource_url}
|
||||
|
||||
# Upstream responded but with non-200 or non-dict payload. For
|
||||
# pass-through servers the gateway is NOT the authorization server,
|
||||
|
|
@ -1661,6 +2242,7 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
|
|||
response_types=data.get("response_types", []),
|
||||
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
|
||||
fallback_client_id=resolved.server_name or resolved.name,
|
||||
client_redirect_uris=data.get("redirect_uris"),
|
||||
)
|
||||
return dummy_return
|
||||
|
||||
|
|
@ -1675,4 +2257,5 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
|
|||
response_types=data.get("response_types", []),
|
||||
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
|
||||
fallback_client_id=mcp_server_name,
|
||||
client_redirect_uris=data.get("redirect_uris"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -73,3 +73,18 @@ class MCPUpstreamAuthError(Exception):
|
|||
detail=detail,
|
||||
headers={"www-authenticate": challenge} if challenge else None,
|
||||
)
|
||||
|
||||
|
||||
class MCPToolResultError(Exception):
|
||||
"""An MCP tool call completed with ``isError=True`` in its result.
|
||||
|
||||
Never raised on the wire path: streamable HTTP MCP correctly returns tool
|
||||
failures as HTTP 200 with ``result.isError: true`` per the MCP spec. This
|
||||
exception only drives the standard failure logging (``status="failure"``
|
||||
payload, OTel ERROR span) for such results.
|
||||
|
||||
Lives here rather than ``utils.py`` deliberately: tests reload ``utils``
|
||||
to re-read its env-derived constants, and a reload would fork this class
|
||||
into two identities, breaking ``isinstance`` checks against instances
|
||||
created before the reload.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -57,7 +57,11 @@ from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
MCP_SAMPLING_AVAILABLE,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
||||
MCPPerUserTokenCache,
|
||||
mcp_per_user_token_cache,
|
||||
resolve_mcp_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
Error,
|
||||
Ok,
|
||||
|
|
@ -70,6 +74,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
|
|||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
InvalidatableOAuthTokenStore,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
|
||||
LazyPerUserOAuthTokenStore,
|
||||
)
|
||||
|
|
@ -78,6 +85,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
AuthorizationCodeConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
|
|
@ -115,7 +123,7 @@ from litellm.proxy.common_utils.user_api_key_cache import get_management_object_
|
|||
from litellm.proxy.utils import ProxyLogging, get_server_root_path
|
||||
from litellm.repositories.table_repositories import MCPServerRepository
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPAuth, MCPStdioConfig
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPStdioConfig
|
||||
from litellm.types.mcp_server.mcp_server_manager import (
|
||||
MCPInfo,
|
||||
MCPOAuthMetadata,
|
||||
|
|
@ -167,6 +175,16 @@ _user_env_vars_cache: dict[tuple[str, str], tuple[dict[str, str], float]] = {}
|
|||
_USER_ENV_VARS_CACHE_TTL = 60 # seconds
|
||||
_USER_ENV_VARS_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
|
||||
|
||||
# Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the
|
||||
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
|
||||
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
|
||||
# config-YAML and DB server loaders so the two paths cannot drift on which modes trigger discovery.
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
|
||||
MCPAuth.oauth2,
|
||||
MCPAuth.true_passthrough,
|
||||
MCPAuth.oauth_delegate,
|
||||
)
|
||||
|
||||
|
||||
def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None:
|
||||
"""Drop a cached entry after the user stores or clears their env var values
|
||||
|
|
@ -213,6 +231,13 @@ def _should_strip_caller_authorization(
|
|||
pass-through cold-start case (RFC 9728) the bearer in
|
||||
``Authorization`` is the upstream OAuth token and must be
|
||||
forwarded, so we keep it.
|
||||
- **oauth_delegate servers**: admission always runs and there is no
|
||||
anonymous path, so the caller's separate ``Authorization`` is
|
||||
forwarded only when a distinct ``x-litellm-api-key`` carried
|
||||
admission. Without that header the ``Authorization`` *was* the
|
||||
admission credential — a virtual key, an IdP JWT, or an SSO / OIDC /
|
||||
session token whose ``api_key`` is ``None`` — and must never reach
|
||||
the upstream, so it is stripped regardless of the ``api_key`` value.
|
||||
"""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
# OBO: the inbound Authorization is the subject token. It is exchanged at the IdP and only the
|
||||
|
|
@ -226,11 +251,13 @@ def _should_strip_caller_authorization(
|
|||
# upstream — it would override another user's stored credential. Delegate and
|
||||
# pass-through return None from to_server_spec and keep forwarding the bearer.
|
||||
return True
|
||||
if not mcp_server.is_oauth_passthrough:
|
||||
if not (mcp_server.is_oauth_passthrough or mcp_server.is_oauth_delegate):
|
||||
return False
|
||||
|
||||
normalized_raw_headers = {str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str)}
|
||||
has_explicit_litellm_admission_header = normalized_raw_headers.get("x-litellm-api-key") is not None
|
||||
if mcp_server.is_oauth_delegate:
|
||||
return not has_explicit_litellm_admission_header
|
||||
admission_consumed_authorization_as_litellm_key = (
|
||||
user_api_key_auth is not None
|
||||
and bool(getattr(user_api_key_auth, "api_key", None))
|
||||
|
|
@ -323,6 +350,89 @@ async def _resolve_byok_mcp_auth_header(
|
|||
return mcp_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: Optional[dict[str, str]],
|
||||
raw_headers: Optional[dict[str, str]],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[dict[str, str]]:
|
||||
"""Egress headers for the client-forwarded-token modes (``true_passthrough`` / ``oauth_delegate``).
|
||||
|
||||
Forwards the caller's ``Authorization`` to the upstream, stripped when
|
||||
``_should_strip_caller_authorization`` says it was consumed as the LiteLLM admission key. Shared by
|
||||
``_call_regular_mcp_tool`` and ``server.py``'s ``_prepare_mcp_server_headers`` so the two egress
|
||||
paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split.
|
||||
"""
|
||||
extra_headers = oauth2_headers.copy() if oauth2_headers else None
|
||||
if extra_headers and _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
return _without_authorization(extra_headers)
|
||||
return extra_headers
|
||||
|
||||
|
||||
def _take_forwarded_authorization(
|
||||
headers: Optional[dict[str, str]],
|
||||
) -> tuple[Optional[str], Optional[dict[str, str]]]:
|
||||
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
|
||||
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
|
||||
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
|
||||
if not headers:
|
||||
return None, headers
|
||||
value = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
|
||||
return value, _without_authorization(headers)
|
||||
|
||||
|
||||
def _passthrough_token_from_mcp_auth_header(
|
||||
mcp_auth_header: Optional[Union[str, dict[str, str]]],
|
||||
) -> Optional[str]:
|
||||
"""The caller's per-server upstream credential for a passthrough-mode server, or None.
|
||||
|
||||
Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated
|
||||
global ``x-mcp-auth`` fallback. Per-server headers are the multi-server shape: they bind one
|
||||
token to one server, so an aggregate scope with several passthrough-mode servers never replays
|
||||
a single credential across upstreams. The value is forwarded verbatim, so it must be the full
|
||||
header value (e.g. ``Bearer <upstream-token>``)."""
|
||||
if isinstance(mcp_auth_header, str):
|
||||
return mcp_auth_header or None
|
||||
if isinstance(mcp_auth_header, dict):
|
||||
return next((v for k, v in mcp_auth_header.items() if k.lower() == "authorization"), None)
|
||||
return None
|
||||
|
||||
|
||||
def _consumes_caller_authorization(server: MCPServer) -> bool:
|
||||
"""True when this server's egress forwards the caller's request-wide ``Authorization`` upstream:
|
||||
the client-forwarded token modes, legacy OAuth pass-through, and legacy upstream-delegated
|
||||
interactive oauth2. An unstamped oauth2 row (flow column not yet backfilled) reads as a consumer,
|
||||
which errs toward suppression — the fail-safe direction."""
|
||||
if server.is_true_passthrough or server.is_oauth_delegate or server.is_oauth_passthrough:
|
||||
return True
|
||||
return (
|
||||
server.auth_type == MCPAuth.oauth2
|
||||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
and not server.has_client_credentials
|
||||
)
|
||||
|
||||
|
||||
def _caller_authorization_fans_out(
|
||||
server: MCPServer,
|
||||
scope_servers: Optional[list[MCPServer]],
|
||||
) -> bool:
|
||||
"""True when forwarding the caller's request-wide ``Authorization`` to ``server`` inside a
|
||||
listing fan-out would replay one credential against multiple upstreams: another server in the
|
||||
scope also consumes it (RFC 9700 cross-resource replay). ``scope_servers`` is None for
|
||||
explicitly-addressed operations (tool call, get_prompt, read_resource, single-server routes),
|
||||
where the client named the one target and the gateway is not choosing recipients."""
|
||||
if scope_servers is None:
|
||||
return False
|
||||
return any(
|
||||
other is not None and other.server_id != server.server_id and _consumes_caller_authorization(other)
|
||||
for other in scope_servers
|
||||
)
|
||||
|
||||
|
||||
def _extract_upstream_auth_failure(
|
||||
exc: BaseException,
|
||||
) -> Optional[tuple[int, Optional[str]]]:
|
||||
|
|
@ -689,9 +799,18 @@ class MCPServerManager:
|
|||
"""
|
||||
return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url)
|
||||
|
||||
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
|
||||
def __init__(
|
||||
self,
|
||||
cred_provider: Optional[UpstreamCredentialProvider] = None,
|
||||
per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None,
|
||||
per_user_token_cache: Optional[MCPPerUserTokenCache] = None,
|
||||
):
|
||||
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
|
||||
self.get_mcp_server_by_id
|
||||
)
|
||||
self._per_user_token_cache = per_user_token_cache or mcp_per_user_token_cache
|
||||
self._cred_provider = cred_provider or UpstreamCredentialProvider(
|
||||
oauth_token_store=LazyPerUserOAuthTokenStore(self.get_mcp_server_by_id),
|
||||
oauth_token_store=self._per_user_oauth_token_store,
|
||||
token_exchanger=build_token_exchanger(),
|
||||
)
|
||||
self.registry: dict[str, MCPServer] = {}
|
||||
|
|
@ -715,8 +834,10 @@ class MCPServerManager:
|
|||
# Per-server outbound tool-call concurrency limiters, lazily created from
|
||||
# each server's max_concurrent_requests. Keyed by server_id so the cap
|
||||
# survives the registry atomic-swap on config reload; a missing key means
|
||||
# the server has no configured limit.
|
||||
self._server_call_semaphores: dict[str, asyncio.Semaphore] = {}
|
||||
# the server has no configured limit. The limit is cached alongside the
|
||||
# semaphore so an edited limit rebuilds it instead of keeping the old cap
|
||||
# until restart.
|
||||
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
|
||||
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
|
||||
"""
|
||||
{
|
||||
|
|
@ -879,7 +1000,7 @@ class MCPServerManager:
|
|||
|
||||
auth_type = server_config.get("auth_type", None)
|
||||
if server_url and (
|
||||
auth_type == MCPAuth.oauth2
|
||||
auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
or self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
server_config.get("token_exchange_endpoint"),
|
||||
|
|
@ -888,7 +1009,7 @@ class MCPServerManager:
|
|||
):
|
||||
mcp_oauth_metadata = await self._descovery_metadata(
|
||||
server_url=server_url,
|
||||
allow_origin_fallback=auth_type == MCPAuth.oauth2,
|
||||
allow_origin_fallback=auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
)
|
||||
else:
|
||||
mcp_oauth_metadata = None
|
||||
|
|
@ -923,6 +1044,24 @@ class MCPServerManager:
|
|||
"browser sign-in, including delegate_auth_to_upstream)."
|
||||
)
|
||||
|
||||
config_dcr_bridge = server_config.get("dcr_bridge", None)
|
||||
if config_dcr_bridge is not None and not isinstance(config_dcr_bridge, bool):
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name or server_id}': dcr_bridge "
|
||||
f"must be a boolean (got {config_dcr_bridge!r})."
|
||||
)
|
||||
if config_dcr_bridge and auth_type not in (
|
||||
MCPAuth.true_passthrough,
|
||||
MCPAuth.oauth_delegate,
|
||||
):
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name or server_id}': dcr_bridge is only "
|
||||
f"supported for auth_type true_passthrough or oauth_delegate (got {auth_type!r}). "
|
||||
"The DCR bridge serves gateway-hosted OAuth discovery for the client-forwarded "
|
||||
"token modes; interactive oauth2 servers already run the gateway "
|
||||
"authorization-code flow."
|
||||
)
|
||||
|
||||
new_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=name_for_prefix,
|
||||
|
|
@ -958,6 +1097,7 @@ class MCPServerManager:
|
|||
available_on_public_internet=bool(server_config.get("available_on_public_internet", True)),
|
||||
delegate_auth_to_upstream=bool(server_config.get("delegate_auth_to_upstream", False)),
|
||||
oauth_passthrough=bool(server_config.get("oauth_passthrough", False)),
|
||||
dcr_bridge=config_dcr_bridge,
|
||||
# AWS SigV4 fields
|
||||
aws_access_key_id=server_config.get("aws_access_key_id", None),
|
||||
aws_secret_access_key=server_config.get("aws_secret_access_key", None),
|
||||
|
|
@ -972,7 +1112,7 @@ class MCPServerManager:
|
|||
audience=server_config.get("audience", None),
|
||||
subject_token_type=server_config.get(
|
||||
"subject_token_type",
|
||||
"urn:ietf:params:oauth:token-type:access_token",
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
),
|
||||
token_exchange_profile=server_config.get("token_exchange_profile", "rfc8693"),
|
||||
allow_sampling=bool(server_config.get("allow_sampling", False)),
|
||||
|
|
@ -1280,17 +1420,18 @@ class MCPServerManager:
|
|||
auth_type = cast(MCPAuthType, mcp_server.auth_type)
|
||||
server_url = mcp_server.url
|
||||
needs_discovery = bool(server_url) and (
|
||||
(auth_type == MCPAuth.oauth2 and not mcp_server.authorization_url)
|
||||
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not mcp_server.authorization_url)
|
||||
or self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None,
|
||||
mcp_server.token_exchange_endpoint
|
||||
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
|
||||
mcp_server.token_url,
|
||||
)
|
||||
)
|
||||
mcp_oauth_metadata = (
|
||||
await self._descovery_metadata(
|
||||
server_url=server_url, # type: ignore[arg-type]
|
||||
allow_origin_fallback=auth_type == MCPAuth.oauth2,
|
||||
allow_origin_fallback=auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
)
|
||||
if needs_discovery
|
||||
else None
|
||||
|
|
@ -1332,6 +1473,7 @@ class MCPServerManager:
|
|||
available_on_public_internet=bool(getattr(mcp_server, "available_on_public_internet", True)),
|
||||
delegate_auth_to_upstream=bool(getattr(mcp_server, "delegate_auth_to_upstream", False)),
|
||||
oauth_passthrough=bool(getattr(mcp_server, "oauth_passthrough", False)),
|
||||
dcr_bridge=getattr(mcp_server, "dcr_bridge", None),
|
||||
created_at=getattr(mcp_server, "created_at", None),
|
||||
updated_at=getattr(mcp_server, "updated_at", None),
|
||||
tool_name_to_display_name=_deserialize_json_dict(getattr(mcp_server, "tool_name_to_display_name", None)),
|
||||
|
|
@ -1349,12 +1491,16 @@ class MCPServerManager:
|
|||
aws_role_name=aws_creds.get("aws_role_name"),
|
||||
aws_session_name=aws_creds.get("aws_session_name"),
|
||||
instructions=mcp_server.instructions,
|
||||
# Token Exchange (OBO) fields — read from credentials JSON blob
|
||||
token_exchange_endpoint=(credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
|
||||
audience=(credentials_dict.get("audience") if credentials_dict else None),
|
||||
subject_token_type=(credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or "urn:ietf:params:oauth:token-type:access_token",
|
||||
token_exchange_profile=(credentials_dict.get("token_exchange_profile") if credentials_dict else None)
|
||||
# Token exchange (OBO) fields: dedicated columns, with the credentials blob as a
|
||||
# back-compat fallback for servers persisted before the columns existed.
|
||||
token_exchange_endpoint=mcp_server.token_exchange_endpoint
|
||||
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
|
||||
audience=mcp_server.audience or (credentials_dict.get("audience") if credentials_dict else None),
|
||||
subject_token_type=mcp_server.subject_token_type
|
||||
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
token_exchange_profile=mcp_server.token_exchange_profile
|
||||
or (credentials_dict.get("token_exchange_profile") if credentials_dict else None)
|
||||
or "rfc8693",
|
||||
timeout=getattr(mcp_server, "timeout", None),
|
||||
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
|
||||
|
|
@ -1618,15 +1764,18 @@ class MCPServerManager:
|
|||
delegate_server_ids = [
|
||||
server.server_id
|
||||
for server in self.get_registry().values()
|
||||
if getattr(server, "auth_type", None) == MCPAuth.oauth2
|
||||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
# M2M servers must not be exposed anonymously: an
|
||||
# unauthenticated caller would get LiteLLM to proxy tool
|
||||
# calls using its stored client_credentials. Resolve the flow
|
||||
# rather than reading has_client_credentials so an unstamped
|
||||
# M2M-shape row (null column, verbatim-read as non-M2M) still
|
||||
# fails closed here, matching the anonymous-delegate auth gate.
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
if (
|
||||
getattr(server, "auth_type", None) == MCPAuth.oauth2
|
||||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
# M2M servers must not be exposed anonymously: an
|
||||
# unauthenticated caller would get LiteLLM to proxy tool
|
||||
# calls using its stored client_credentials. Resolve the flow
|
||||
# rather than reading has_client_credentials so an unstamped
|
||||
# M2M-shape row (null column, verbatim-read as non-M2M) still
|
||||
# fails closed here, matching the anonymous-delegate auth gate.
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
)
|
||||
or getattr(server, "auth_type", None) == MCPAuth.true_passthrough
|
||||
]
|
||||
combined_servers.update(delegate_server_ids)
|
||||
|
||||
|
|
@ -2224,16 +2373,17 @@ class MCPServerManager:
|
|||
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
|
||||
provider = cred_provider or self._cred_provider
|
||||
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
|
||||
# so it wins - except for the per-user modes the v2 resolver owns (authorization_code's
|
||||
# stored token and token_exchange's RFC 8693 minted token). A caller must not be able to
|
||||
# substitute another user's stored credential, nor silently disable the OBO exchange and
|
||||
# forward an arbitrary bearer upstream, so we keep the v2 spec and ignore the override for
|
||||
# both; the REST tools preview supplies its not-yet-persisted token through the resolver
|
||||
# (cred_provider), never this path.
|
||||
# so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's
|
||||
# stored token, token_exchange's RFC 8693 minted token, and the passthrough modes'
|
||||
# forwarded caller token). A caller must not be able to substitute another user's stored
|
||||
# credential, nor silently disable the OBO exchange and forward an arbitrary bearer
|
||||
# upstream, so we keep the v2 spec and ignore the override for these; the REST tools
|
||||
# preview supplies its not-yet-persisted token through the resolver (cred_provider),
|
||||
# never this path.
|
||||
if (
|
||||
spec is not None
|
||||
and mcp_auth_header
|
||||
and not isinstance(spec.config, (AuthorizationCodeConfig, TokenExchangeConfig))
|
||||
and not isinstance(spec.config, (AuthorizationCodeConfig, PassthroughConfig, TokenExchangeConfig))
|
||||
):
|
||||
spec = None
|
||||
auth_value = (
|
||||
|
|
@ -2300,11 +2450,17 @@ class MCPServerManager:
|
|||
server_url = server.url or ""
|
||||
|
||||
if spec is not None:
|
||||
inbound_token = subject_token
|
||||
if isinstance(spec.config, PassthroughConfig):
|
||||
inbound_token, extra_headers = _take_forwarded_authorization(extra_headers)
|
||||
per_server_token = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
if per_server_token is not None:
|
||||
inbound_token = per_server_token
|
||||
resolved_auth, extra_headers = await self._resolve_v2_auth(
|
||||
server=server,
|
||||
spec=spec,
|
||||
provider=provider,
|
||||
subject_token=subject_token,
|
||||
subject_token=inbound_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
|
|
@ -2474,10 +2630,16 @@ class MCPServerManager:
|
|||
|
||||
return prefixed_or_original_tools
|
||||
|
||||
except MCPUpstreamAuthError:
|
||||
except MCPUpstreamAuthError as upstream_auth_error:
|
||||
# Pass-through 401 must surface to single-server routes so the
|
||||
# client triggers the upstream OAuth flow. The multi-server
|
||||
# aggregator catches this explicitly to keep absorbing.
|
||||
if server.is_dcr_bridge and upstream_auth_error.www_authenticate is not None:
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=upstream_auth_error.status_code,
|
||||
www_authenticate=None,
|
||||
server_name=upstream_auth_error.server_name,
|
||||
) from upstream_auth_error
|
||||
raise
|
||||
except HTTPException as e:
|
||||
# A v2 resolver auth challenge (token_exchange's RFC 9728 401, authorization_code's
|
||||
|
|
@ -2487,9 +2649,10 @@ class MCPServerManager:
|
|||
# Non-auth HTTP errors stay absorbed so one misconfigured server can't blank the listing.
|
||||
if e.status_code in (401, 403):
|
||||
headers = e.headers or {}
|
||||
challenge_header = headers.get("WWW-Authenticate") or headers.get("www-authenticate")
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=e.status_code,
|
||||
www_authenticate=headers.get("WWW-Authenticate") or headers.get("www-authenticate"),
|
||||
www_authenticate=None if server.is_dcr_bridge else challenge_header,
|
||||
server_name=server.name,
|
||||
) from e
|
||||
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
|
||||
|
|
@ -3589,10 +3752,11 @@ class MCPServerManager:
|
|||
limit = mcp_server.max_concurrent_requests
|
||||
if limit is None or limit <= 0:
|
||||
return None
|
||||
semaphore = self._server_call_semaphores.get(mcp_server.server_id)
|
||||
if semaphore is None:
|
||||
semaphore = asyncio.Semaphore(limit)
|
||||
self._server_call_semaphores[mcp_server.server_id] = semaphore
|
||||
cached = self._server_call_semaphores.get(mcp_server.server_id)
|
||||
if cached is not None and cached[0] == limit:
|
||||
return cached[1]
|
||||
semaphore = asyncio.Semaphore(limit)
|
||||
self._server_call_semaphores[mcp_server.server_id] = (limit, semaphore)
|
||||
return semaphore
|
||||
|
||||
@asynccontextmanager
|
||||
|
|
@ -3725,6 +3889,13 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
elif mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate:
|
||||
extra_headers = _client_forwarded_authorization_headers(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if mcp_server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
@ -3806,22 +3977,66 @@ class MCPServerManager:
|
|||
# OBO: the exchanged token may have been revoked/rotated upstream since it was cached, so
|
||||
# an upstream 401 gets one re-mint + retry. Gated to this mode; all others keep the plain
|
||||
# single call below.
|
||||
tool_call_coro = self._obo_call_tool_with_retry(
|
||||
client=client,
|
||||
call_tool_params=call_tool_params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
mcp_server=mcp_server,
|
||||
server_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
async def _obo_call_tool_limited():
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._obo_call_tool_with_retry(
|
||||
client=client,
|
||||
call_tool_params=call_tool_params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
mcp_server=mcp_server,
|
||||
server_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
tool_call_coro = _obo_call_tool_limited()
|
||||
else:
|
||||
# Scoped to the two client-forwarded token modes this stack introduced; legacy
|
||||
# oauth2 + delegate_auth_to_upstream (is_oauth_passthrough) is being removed, so it is not
|
||||
# added here even though the list path still relays for it.
|
||||
relays_upstream_auth = mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate
|
||||
server_label = mcp_server.name or mcp_server.server_name or mcp_server.alias or ""
|
||||
|
||||
async def _call_tool_via_client(client, params):
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
if not relays_upstream_auth:
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
# The client-forwarded modes carry the caller's own upstream token, so an upstream
|
||||
# 401 (expired/invalid token) is the caller's to resolve: relay it as
|
||||
# MCPUpstreamAuthError so single-server REST callers turn it into a 401 +
|
||||
# WWW-Authenticate and re-run the upstream OAuth flow. Only 401 is a re-auth signal
|
||||
# (mirrors the list path and MCPUpstreamAuthError's contract); a 403 is a genuine
|
||||
# authorization failure that re-auth won't fix, so it takes the non-auth branch and
|
||||
# stays a visible warning. raise_on_error only re-raises transport failures
|
||||
# (tool-level isError results are still returned normally); a non-auth failure keeps
|
||||
# the same isError degradation the default path produces.
|
||||
try:
|
||||
return await client.call_tool(
|
||||
params, host_progress_callback=host_progress_callback, raise_on_error=True
|
||||
)
|
||||
except Exception as e:
|
||||
auth_info = _extract_upstream_auth_failure(e)
|
||||
if auth_info is None or auth_info[0] != 401:
|
||||
# A genuine (non-auth or 403-forbidden) upstream/transport failure.
|
||||
# raise_on_error demoted the client-layer log to debug, so surface it here at
|
||||
# warning level to keep the outage visible; the caller still gets the graceful
|
||||
# isError result the default masking path would have produced. Log the
|
||||
# exception type only, never str(e), which for an httpx error embeds the
|
||||
# upstream URL (a credential can hide in it).
|
||||
verbose_logger.warning(
|
||||
"Pass-through MCP tool call failed against %s (non-auth, %s)",
|
||||
server_label,
|
||||
type(e).__name__,
|
||||
)
|
||||
return client.error_tool_result(e)
|
||||
_, www_authenticate = auth_info
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=401,
|
||||
www_authenticate=www_authenticate,
|
||||
server_name=server_label,
|
||||
) from e
|
||||
|
||||
tool_call_coro = _call_tool_via_client(client, call_tool_params)
|
||||
|
||||
|
|
@ -3910,6 +4125,28 @@ class MCPServerManager:
|
|||
return False
|
||||
return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec)
|
||||
|
||||
async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None:
|
||||
"""Drop every cached token for ``(user_id, server_id)`` after the credential row changes
|
||||
(re-auth, revoke, config-change purge): the v2 chain's cache and the legacy per-user token
|
||||
cache, so the next resolve reads the new row instead of serving the replaced token until its
|
||||
cache TTL, whichever path resolves it. This is the single invalidation point for per-user
|
||||
OAuth tokens; callers must not evict individual caches directly. Best-effort: a cache-drop
|
||||
failure is logged, never raised, because the DB write already succeeded and the TTL remains
|
||||
the backstop.
|
||||
"""
|
||||
try:
|
||||
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
|
||||
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
|
||||
verbose_logger.warning(
|
||||
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
try:
|
||||
await self._per_user_token_cache.delete(user_id, server_id)
|
||||
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
|
||||
verbose_logger.warning(
|
||||
"Failed to drop legacy cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
|
||||
async def _resolve_oauth2_headers_for_tool_call(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -4630,6 +4867,11 @@ class MCPServerManager:
|
|||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
oauth2_flow=server.oauth2_flow,
|
||||
dcr_bridge=server.dcr_bridge,
|
||||
token_exchange_endpoint=server.token_exchange_endpoint,
|
||||
audience=server.audience,
|
||||
subject_token_type=server.subject_token_type,
|
||||
token_exchange_profile=server.token_exchange_profile,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
|
|
@ -4734,10 +4976,15 @@ class MCPServerManager:
|
|||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
oauth2_flow=server.oauth2_flow,
|
||||
token_exchange_endpoint=server.token_exchange_endpoint,
|
||||
audience=server.audience,
|
||||
subject_token_type=server.subject_token_type,
|
||||
token_exchange_profile=server.token_exchange_profile,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
available_on_public_internet=server.available_on_public_internet,
|
||||
delegate_auth_to_upstream=server.delegate_auth_to_upstream,
|
||||
oauth_passthrough=getattr(server, "oauth_passthrough", False),
|
||||
dcr_bridge=server.dcr_bridge,
|
||||
is_byok=server.is_byok,
|
||||
byok_description=server.byok_description,
|
||||
byok_api_key_help_url=server.byok_api_key_help_url,
|
||||
|
|
|
|||
|
|
@ -129,7 +129,7 @@ def get_request_base_url(request: Request) -> str:
|
|||
if x_forwarded_port and ":" not in netloc:
|
||||
netloc = f"{netloc}:{x_forwarded_port}"
|
||||
|
||||
return urlunparse((scheme, netloc, parsed.path, "", "", ""))
|
||||
return urlunparse((scheme, _strip_default_port(scheme, netloc), parsed.path, "", "", ""))
|
||||
|
||||
|
||||
def validate_loopback_redirect_uri(redirect_uri: str) -> None:
|
||||
|
|
|
|||
|
|
@ -23,12 +23,13 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
AuthorizationCodeConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
Subject,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -62,9 +63,10 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
|
||||
explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live
|
||||
modes: ``none``, the static-header family (``api_key`` plus the Authorization schemes,
|
||||
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), and
|
||||
``oauth2_token_exchange`` (OBO); client_credentials (M2M), delegated/passthrough
|
||||
oauth2, and SigV4 return None and stay on v1.
|
||||
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2_token_exchange``
|
||||
(OBO), and the client-forwarded token modes ``true_passthrough`` / ``oauth_delegate``
|
||||
(``PassthroughConfig``); client_credentials (M2M), delegated/passthrough oauth2, and SigV4
|
||||
return None and stay on v1.
|
||||
"""
|
||||
if server.is_byok:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
|
||||
|
|
@ -94,6 +96,8 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
)
|
||||
# client_credentials (M2M) and delegate/passthrough oauth2 stay on v1
|
||||
return None
|
||||
case MCPAuth.true_passthrough | MCPAuth.oauth_delegate:
|
||||
return ServerSpec(server_id=server.server_id, resource=resource, config=PassthroughConfig())
|
||||
case MCPAuth.oauth2_token_exchange:
|
||||
return _token_exchange_spec(server, resource)
|
||||
case MCPAuth.aws_sigv4:
|
||||
|
|
@ -124,7 +128,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpe
|
|||
resource=resource,
|
||||
config=TokenExchangeConfig(
|
||||
profile=profile,
|
||||
subject_token_type=server.subject_token_type or "urn:ietf:params:oauth:token-type:access_token",
|
||||
subject_token_type=server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
token_exchange_endpoint=endpoint,
|
||||
audience=server.audience,
|
||||
client_id=server.client_id,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,169 @@
|
|||
"""Producer and consumer helpers for the DCR-bridge ``oauth_delegate`` envelope.
|
||||
|
||||
A DCR-bridge ``oauth_delegate`` client presents ONE bearer that is a litellm-signed
|
||||
envelope (see :mod:`.envelope`) carrying both a litellm identity and the upstream OAuth
|
||||
token. The gateway token endpoint mints it (producer) at OAuth issuance, and at the MCP
|
||||
admission edge the gateway derives the envelope keys from the proxy ``master_key``, opens
|
||||
it, admits the request under the recovered identity, and forwards the inner upstream token
|
||||
to the upstream MCP server (consumer). This module is the pure surface for both sides; the
|
||||
token-endpoint and admission wiring live in their respective call sites.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from typing import Literal, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, SecretStr
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
EnvelopeIdentity,
|
||||
EnvelopeKeys,
|
||||
EnvelopeMintError,
|
||||
OpenedEnvelope,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
is_envelope,
|
||||
mint_envelope,
|
||||
open_envelope,
|
||||
)
|
||||
|
||||
_SIGNING_KEY_DOMAIN = b"litellm-mcp-bridge:envelope-signing:"
|
||||
_ENCRYPTION_KEY_DOMAIN = b"litellm-mcp-bridge:envelope-encryption:"
|
||||
|
||||
# scrypt work factors (RFC 7914). n=2**15 with r=8/p=1 costs ~50ms and ~32MB per derivation, which
|
||||
# makes offline guessing of a candidate master key memory-hard rather than a bare hash comparison.
|
||||
_SCRYPT_N = 2**15
|
||||
_SCRYPT_R = 8
|
||||
_SCRYPT_P = 1
|
||||
# scrypt's working-set is ~128 * N * r * p bytes; cap at twice that so the maxmem ceiling scales
|
||||
# with every work factor and a future p or r bump does not trip "memory limit exceeded".
|
||||
_SCRYPT_MAXMEM = 128 * _SCRYPT_N * _SCRYPT_R * _SCRYPT_P * 2
|
||||
_DERIVED_KEY_BYTES = 32
|
||||
|
||||
|
||||
@lru_cache(maxsize=8)
|
||||
def envelope_keys_from_master_key(master_key: str) -> EnvelopeKeys:
|
||||
"""Derive the envelope signing and encryption keys from the proxy master key.
|
||||
|
||||
A memory-hard scrypt KDF (RFC 7914) over two distinct domain-label salts yields two
|
||||
independent 256-bit subkeys from the one secret, so the producer (mint) and consumer
|
||||
(open) agree on keys without persisting any. scrypt is used rather than a bare hash or
|
||||
HMAC so that a captured envelope is not a cheap offline oracle for the master key: each
|
||||
candidate guess costs a full memory-hard derivation, which is what protects a deployment
|
||||
whose master key is weaker than it should be. The result is cached (the master key is
|
||||
fixed for a process), so the KDF runs once per key and adds nothing to the per-request
|
||||
admission path. The derivation is deterministic; rotating ``master_key`` invalidates
|
||||
every outstanding envelope, which is the intended behavior for a signing-key change.
|
||||
"""
|
||||
signing = hashlib.scrypt(
|
||||
master_key.encode(),
|
||||
salt=_SIGNING_KEY_DOMAIN,
|
||||
n=_SCRYPT_N,
|
||||
r=_SCRYPT_R,
|
||||
p=_SCRYPT_P,
|
||||
maxmem=_SCRYPT_MAXMEM,
|
||||
dklen=_DERIVED_KEY_BYTES,
|
||||
).hex()
|
||||
encryption = hashlib.scrypt(
|
||||
master_key.encode(),
|
||||
salt=_ENCRYPTION_KEY_DOMAIN,
|
||||
n=_SCRYPT_N,
|
||||
r=_SCRYPT_R,
|
||||
p=_SCRYPT_P,
|
||||
maxmem=_SCRYPT_MAXMEM,
|
||||
dklen=_DERIVED_KEY_BYTES,
|
||||
).hex()
|
||||
return EnvelopeKeys(signing_key=SecretStr(signing), encryption_key=SecretStr(encryption))
|
||||
|
||||
|
||||
def build_bridge_token_response(
|
||||
identity: EnvelopeIdentity,
|
||||
grant: UpstreamTokenGrant,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> SealedEnvelope | EnvelopeMintError:
|
||||
"""Seal ``grant`` for ``identity`` into the client-held bearer the token endpoint returns.
|
||||
|
||||
The producer mirror of :func:`resolve_bridge_envelope`: a thin, pure wrapper over
|
||||
:func:`mint_envelope` that returns the sealed envelope, or the mint error as a value
|
||||
(an oversized grant) for the caller to map onto an OAuth error response.
|
||||
"""
|
||||
return mint_envelope(identity, grant, keys, now)
|
||||
|
||||
|
||||
class NotBridgeEnvelope(BaseModel):
|
||||
"""The bearer is not an envelope; admission continues on its normal path."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["not_bridge_envelope"] = "not_bridge_envelope"
|
||||
|
||||
|
||||
class BridgeEnvelopeAdmitted(BaseModel):
|
||||
"""A valid envelope: the identity to admit under and the full upstream ``Authorization``
|
||||
value (``token_type access_token``) to forward to the upstream MCP server."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["admitted"] = "admitted"
|
||||
identity: EnvelopeIdentity
|
||||
upstream_authorization: SecretStr
|
||||
|
||||
|
||||
class BridgeEnvelopeInvalid(BaseModel):
|
||||
"""The bearer is envelope-shaped but did not open (expired, tampered, wrong key);
|
||||
admission must fail closed rather than fall through to normal validation."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["invalid"] = "invalid"
|
||||
|
||||
|
||||
BridgeEnvelopeResult: TypeAlias = NotBridgeEnvelope | BridgeEnvelopeAdmitted | BridgeEnvelopeInvalid
|
||||
|
||||
|
||||
def _strip_bearer(value: str) -> str:
|
||||
parts = value.split(None, 1)
|
||||
if len(parts) == 2 and parts[0].lower() == "bearer":
|
||||
return parts[1]
|
||||
return value
|
||||
|
||||
|
||||
def is_bridge_envelope_shaped(authorization_value: str) -> bool:
|
||||
"""Cheap, keyless test that an ``Authorization`` value carries an envelope (optional
|
||||
``Bearer`` scheme stripped). The admission edge engages the bridge arm only for an
|
||||
envelope, so a plain upstream bearer falls through to normal oauth2 admission."""
|
||||
return is_envelope(_strip_bearer(authorization_value))
|
||||
|
||||
|
||||
def resolve_bridge_envelope(
|
||||
authorization_value: str,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
expected_server_id: str,
|
||||
) -> BridgeEnvelopeResult:
|
||||
"""Classify an ``Authorization`` value presented to a bridge ``oauth_delegate`` server.
|
||||
|
||||
Strips an optional ``Bearer`` scheme, then returns ``NotBridgeEnvelope`` for a
|
||||
non-envelope bearer (normal admission continues), ``BridgeEnvelopeAdmitted`` with the
|
||||
recovered identity and the upstream ``Authorization`` value to forward for a valid
|
||||
envelope, and ``BridgeEnvelopeInvalid`` for an envelope-shaped bearer that will not
|
||||
open. Never raises: it is total over hostile input via :func:`open_envelope`.
|
||||
|
||||
``expected_server_id`` is the ``server_id`` of the MCP server the request targets; an
|
||||
opened envelope whose sealed ``server_id`` does not match is rejected as
|
||||
``BridgeEnvelopeInvalid``. Binding here (rather than leaving it to the caller) prevents
|
||||
replaying an envelope minted for one server against another, which would forward the
|
||||
first server's upstream credential across a server boundary. ``server_id`` is not a
|
||||
secret (the caller targets that server), so a plain equality check is sufficient and,
|
||||
unlike ``hmac.compare_digest`` on ``str``, does not raise on a non-ASCII server_id.
|
||||
"""
|
||||
candidate = _strip_bearer(authorization_value)
|
||||
if not is_envelope(candidate):
|
||||
return NotBridgeEnvelope()
|
||||
opened = open_envelope(candidate, keys, now)
|
||||
if not isinstance(opened, OpenedEnvelope):
|
||||
return BridgeEnvelopeInvalid()
|
||||
if opened.identity.server_id != expected_server_id:
|
||||
return BridgeEnvelopeInvalid()
|
||||
grant = opened.grant
|
||||
upstream_authorization = f"{grant.token_type} {grant.access_token.get_secret_value()}"
|
||||
return BridgeEnvelopeAdmitted(identity=opened.identity, upstream_authorization=SecretStr(upstream_authorization))
|
||||
|
|
@ -0,0 +1,366 @@
|
|||
"""Client-held sealed envelope for the oauth_delegate DCR bridge.
|
||||
|
||||
A DCR-bridge client holds ONE bearer that must carry BOTH a litellm identity and the
|
||||
upstream OAuth grant, with zero server-side storage. The gateway token endpoint mints a
|
||||
litellm-signed envelope (:func:`mint_envelope`); the MCP edge validates it, recovers the
|
||||
identity claims and the inner upstream grant (:func:`open_envelope`), and forwards the
|
||||
inner access token upstream. This module is pure and unwired: it imports nothing from
|
||||
endpoint or edge code, reads no proxy globals, and takes all key material and the clock
|
||||
as explicit parameters.
|
||||
|
||||
Wire shape: ``llm_env_`` + an HS256 JWT (same signing approach as the BYOK session
|
||||
bearer in ``byok_oauth_endpoints.py``). Registered claims are ``iss``/``iat``/``exp``;
|
||||
custom claims are ``server_id``, ``key_hash``, and ``grant``, where ``grant`` is the
|
||||
upstream token grant serialized to JSON, encrypted with the repo's symmetric
|
||||
encryption helpers (``encrypt_value``/``decrypt_value`` from
|
||||
``encrypt_decrypt_utils`` — the same family ``encrypt_value_helper`` applies to
|
||||
persisted DCR credentials), and base64url-encoded, so the inner token never appears
|
||||
in plaintext anywhere in the envelope.
|
||||
|
||||
Failures are values: :func:`open_envelope` returns one of the frozen
|
||||
``EnvelopeOpenError`` variants (discriminated on ``tag``) for invalid, expired,
|
||||
tampered, or undecryptable input, and :func:`mint_envelope` returns
|
||||
``EnvelopeTooLarge`` for oversized grants. Error values carry tags and sizes only,
|
||||
never token material.
|
||||
|
||||
The pydantic input models reject programmer errors at construction (e.g. a
|
||||
non-positive ``expires_in`` or an empty required field). :func:`open_envelope` is
|
||||
additionally total over hostile, attacker-controlled input: it never raises, only
|
||||
returns an ``EnvelopeOpenError``. :func:`mint_envelope` operates on a
|
||||
gateway-supplied grant (an upstream IdP's UTF-8 JSON token response), so it does not
|
||||
defend against non-UTF-8 field content that cannot survive JSON parsing; its only
|
||||
value-typed failure is ``EnvelopeTooLarge``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Literal, TypeAlias
|
||||
|
||||
import jwt
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value, encrypt_value
|
||||
|
||||
ENVELOPE_PREFIX = "llm_env_"
|
||||
"""Marker prefix on every serialized envelope so the edge can cheaply tell an envelope
|
||||
from a raw upstream token before doing any cryptography."""
|
||||
|
||||
ENVELOPE_ISSUER = "litellm-mcp-bridge"
|
||||
"""``iss`` claim stamped into every envelope and required back on open."""
|
||||
|
||||
MAX_ENVELOPE_TTL_SECONDS = 3600
|
||||
"""Hard ceiling on envelope lifetime. ``exp`` is ``min(upstream expires_in, this cap)``
|
||||
(the cap alone when the upstream omits ``expires_in``), matching the 1h lifetime of the
|
||||
BYOK session bearer this module's signing approach is borrowed from: a client-held
|
||||
credential should never outlive a bounded window even when the upstream token does."""
|
||||
|
||||
MAX_ENVELOPE_BYTES = 12288
|
||||
"""Size cap on the final serialized envelope (prefix + JWT, in bytes). Upstream JWTs
|
||||
commonly run 2-4KB; base64 plus encryption overhead roughly doubles that inside the
|
||||
envelope, and common proxy/server header limits sit around 16KB total. 12288 leaves
|
||||
comfortable headroom for a large upstream token while keeping the envelope safely
|
||||
transmittable as a single Authorization header. Oversized grants are rejected with a
|
||||
typed error, never truncated."""
|
||||
|
||||
_ENVELOPE_JWT_ALGORITHM = "HS256"
|
||||
|
||||
|
||||
class EnvelopeIdentity(BaseModel):
|
||||
"""The litellm identity the envelope binds the inner grant to.
|
||||
|
||||
``key_hash`` is the hashed litellm key that authorized the mint, never a raw
|
||||
credential (and the edge rejects a bare hash presented as a bearer). Admission
|
||||
reloads the live key record by it, so the key's current team/org/object-permission
|
||||
restrictions and its revocation state are enforced at use time rather than frozen at
|
||||
mint time. ``server_id`` binds the envelope to one MCP server so it cannot be replayed
|
||||
across a server boundary.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
server_id: str = Field(min_length=1)
|
||||
key_hash: str = Field(min_length=1)
|
||||
|
||||
|
||||
class UpstreamTokenGrant(BaseModel):
|
||||
"""The upstream OAuth token response fields sealed inside the envelope.
|
||||
|
||||
``expires_in`` must be positive when present; a non-positive value is a programmer
|
||||
error rejected at construction. Token fields are ``SecretStr`` so reprs never leak
|
||||
them.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
access_token: SecretStr = Field(min_length=1)
|
||||
token_type: str = Field(min_length=1)
|
||||
refresh_token: SecretStr | None = None
|
||||
scope: str | None = None
|
||||
expires_in: int | None = Field(default=None, gt=0)
|
||||
|
||||
|
||||
class EnvelopeKeys(BaseModel):
|
||||
"""Injected key material: the HS256 signing key and the symmetric encryption key.
|
||||
|
||||
``signing_key`` must be at least 32 bytes: HS256's HMAC-SHA256 has a 256-bit
|
||||
security level, RFC 7518 requires a key of at least that size, and a shorter key
|
||||
makes PyJWT emit ``InsecureKeyLengthWarning``.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
signing_key: SecretStr = Field(min_length=32)
|
||||
encryption_key: SecretStr = Field(min_length=1)
|
||||
|
||||
|
||||
class SealedEnvelope(BaseModel):
|
||||
"""A minted envelope: the client-held bearer value and when it expires."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
token: SecretStr
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class OpenedEnvelope(BaseModel):
|
||||
"""A validated envelope: the identity it was minted for and the recovered grant."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
identity: EnvelopeIdentity
|
||||
grant: UpstreamTokenGrant
|
||||
|
||||
|
||||
class EnvelopeTooLarge(BaseModel):
|
||||
"""The serialized envelope exceeded ``MAX_ENVELOPE_BYTES``; carries sizes only."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["envelope_too_large"] = "envelope_too_large"
|
||||
size_bytes: int
|
||||
max_bytes: int
|
||||
|
||||
|
||||
EnvelopeMintError: TypeAlias = EnvelopeTooLarge
|
||||
|
||||
|
||||
class NotAnEnvelope(BaseModel):
|
||||
"""The candidate does not carry the envelope prefix."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["not_an_envelope"] = "not_an_envelope"
|
||||
|
||||
|
||||
class BadSignature(BaseModel):
|
||||
"""The JWT signature does not verify under the provided signing key."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["bad_signature"] = "bad_signature"
|
||||
|
||||
|
||||
class Expired(BaseModel):
|
||||
"""The envelope's ``exp`` is not in the future relative to the provided ``now``."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["expired"] = "expired"
|
||||
|
||||
|
||||
class MalformedPayload(BaseModel):
|
||||
"""The token is not a well-formed envelope: undecodable JWT, wrong issuer, missing
|
||||
or mistyped claims, or a decrypted grant that fails validation."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["malformed_payload"] = "malformed_payload"
|
||||
|
||||
|
||||
class DecryptFailed(BaseModel):
|
||||
"""The signed ``grant`` blob could not be decrypted under the provided key."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["decrypt_failed"] = "decrypt_failed"
|
||||
|
||||
|
||||
EnvelopeOpenError: TypeAlias = NotAnEnvelope | BadSignature | Expired | MalformedPayload | DecryptFailed
|
||||
|
||||
|
||||
class _EnvelopeClaims(BaseModel):
|
||||
"""Decoded-claims boundary that pins the exact shape :func:`mint_envelope` emits.
|
||||
|
||||
``server_id``/``key_hash`` mirror the ``min_length`` constraints of
|
||||
:class:`EnvelopeIdentity` so any claim set that validates here also constructs an
|
||||
identity, keeping :func:`open_envelope` raise-free: a correctly signed JWT with an
|
||||
empty identity claim fails here and maps to ``MalformedPayload``.
|
||||
|
||||
``strict`` rejects coerced types (``exp: "123"``, ``exp: 123.0``) rather than opening
|
||||
on them, and ``extra="forbid"`` rejects any claim the gateway never mints (a hostile
|
||||
``nbf``/``aud``/... rides along on a re-signed token). Since PyJWT's own ``iat``/
|
||||
``nbf``/``exp`` validators are disabled at decode (they raise on hostile claim types
|
||||
and, for ``iat``/``nbf``, compare against the wall clock rather than the injected
|
||||
``now``), this model is the sole, total type gate for every registered claim.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
|
||||
iss: str
|
||||
iat: int
|
||||
exp: int
|
||||
server_id: str = Field(min_length=1)
|
||||
key_hash: str = Field(min_length=1)
|
||||
grant: str = Field(min_length=1)
|
||||
|
||||
|
||||
class _GrantWire(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
access_token: str
|
||||
token_type: str
|
||||
refresh_token: str | None = None
|
||||
scope: str | None = None
|
||||
expires_in: int | None = None
|
||||
|
||||
|
||||
def is_envelope(candidate: str) -> bool:
|
||||
"""Cheap prefix check so the edge can route envelopes vs raw tokens without crypto."""
|
||||
return candidate.startswith(ENVELOPE_PREFIX)
|
||||
|
||||
|
||||
def mint_envelope(
|
||||
identity: EnvelopeIdentity,
|
||||
grant: UpstreamTokenGrant,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> SealedEnvelope | EnvelopeMintError:
|
||||
"""Seal ``grant`` for ``identity`` into a client-held envelope.
|
||||
|
||||
``exp`` is ``min(grant.expires_in, MAX_ENVELOPE_TTL_SECONDS)`` seconds from ``now``
|
||||
(the cap alone when ``expires_in`` is absent). Returns ``EnvelopeTooLarge`` when the
|
||||
serialized envelope exceeds ``MAX_ENVELOPE_BYTES``.
|
||||
"""
|
||||
expires_at = now + timedelta(seconds=_envelope_ttl_seconds(grant.expires_in))
|
||||
claims = _EnvelopeClaims(
|
||||
iss=ENVELOPE_ISSUER,
|
||||
iat=int(now.timestamp()),
|
||||
exp=int(expires_at.timestamp()),
|
||||
server_id=identity.server_id,
|
||||
key_hash=identity.key_hash,
|
||||
grant=_encrypt_grant_blob(_grant_plaintext(grant), keys.encryption_key),
|
||||
)
|
||||
token = ENVELOPE_PREFIX + jwt.encode(
|
||||
claims.model_dump(),
|
||||
keys.signing_key.get_secret_value(),
|
||||
algorithm=_ENVELOPE_JWT_ALGORITHM,
|
||||
)
|
||||
size_bytes = len(token.encode("utf-8"))
|
||||
if size_bytes > MAX_ENVELOPE_BYTES:
|
||||
return EnvelopeTooLarge(size_bytes=size_bytes, max_bytes=MAX_ENVELOPE_BYTES)
|
||||
return SealedEnvelope(token=SecretStr(token), expires_at=expires_at)
|
||||
|
||||
|
||||
def open_envelope(
|
||||
candidate: str,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> OpenedEnvelope | EnvelopeOpenError:
|
||||
"""Validate ``candidate`` and recover the identity and inner grant.
|
||||
|
||||
Never raises for bad input: every invalid, expired, tampered, or undecryptable
|
||||
candidate maps to a distinct ``EnvelopeOpenError`` variant. The recovered
|
||||
``grant.expires_in`` is the value the upstream reported at mint time and is not
|
||||
re-derived, so it is stale by up to the envelope's lifetime; callers that need a
|
||||
live remaining lifetime should use ``now`` against the upstream, not this field.
|
||||
"""
|
||||
if not is_envelope(candidate):
|
||||
return NotAnEnvelope()
|
||||
# UTF-8 byte length is never below character length, so a character count already over the
|
||||
# cap rejects an oversize candidate in O(1) without encoding it; the exact byte check then
|
||||
# runs only on candidates already bounded to <= MAX_ENVELOPE_BYTES characters.
|
||||
if len(candidate) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
if len(candidate.encode("utf-8", "surrogatepass")) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
claims = _decode_claims(candidate.removeprefix(ENVELOPE_PREFIX), keys.signing_key)
|
||||
if not isinstance(claims, _EnvelopeClaims):
|
||||
return claims
|
||||
if now.timestamp() >= claims.exp:
|
||||
return Expired()
|
||||
grant = _decrypt_grant(claims.grant, keys.encryption_key)
|
||||
if not isinstance(grant, UpstreamTokenGrant):
|
||||
return grant
|
||||
return OpenedEnvelope(
|
||||
identity=EnvelopeIdentity(server_id=claims.server_id, key_hash=claims.key_hash),
|
||||
grant=grant,
|
||||
)
|
||||
|
||||
|
||||
def _envelope_ttl_seconds(upstream_expires_in: int | None) -> int:
|
||||
if upstream_expires_in is None:
|
||||
return MAX_ENVELOPE_TTL_SECONDS
|
||||
return min(upstream_expires_in, MAX_ENVELOPE_TTL_SECONDS)
|
||||
|
||||
|
||||
def _grant_plaintext(grant: UpstreamTokenGrant) -> str:
|
||||
wire = _GrantWire(
|
||||
access_token=grant.access_token.get_secret_value(),
|
||||
token_type=grant.token_type,
|
||||
refresh_token=None if grant.refresh_token is None else grant.refresh_token.get_secret_value(),
|
||||
scope=grant.scope,
|
||||
expires_in=grant.expires_in,
|
||||
)
|
||||
return wire.model_dump_json(exclude_none=True)
|
||||
|
||||
|
||||
def _decode_claims(
|
||||
compact: str,
|
||||
signing_key: SecretStr,
|
||||
) -> _EnvelopeClaims | BadSignature | MalformedPayload:
|
||||
"""Verify the HS256 signature and shape of an attacker-controlled compact JWT.
|
||||
|
||||
``compact`` is fully hostile and bounded to ``MAX_ENVELOPE_BYTES`` by the caller.
|
||||
PyJWT's ``iat``/``nbf``/``exp`` validators are disabled: they raise on hostile claim
|
||||
types and, for ``iat``/``nbf``, compare against the wall clock rather than the
|
||||
injected ``now`` (``exp`` is checked by the caller against ``now``). Apart from a
|
||||
signature mismatch (``BadSignature``), every decode failure is ``MalformedPayload``:
|
||||
a non-UTF-8 candidate surfaces as ``UnicodeEncodeError`` (a ``ValueError``), a
|
||||
non-string registered claim such as ``iss`` as a ``TypeError`` from PyJWT's claim
|
||||
validators, and a wrong issuer or structurally invalid token as an
|
||||
``InvalidTokenError``. ``_EnvelopeClaims`` is the total type gate for the payload.
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
compact,
|
||||
signing_key.get_secret_value(),
|
||||
algorithms=[_ENVELOPE_JWT_ALGORITHM],
|
||||
issuer=ENVELOPE_ISSUER,
|
||||
options={
|
||||
"verify_exp": False,
|
||||
"verify_iat": False,
|
||||
"verify_nbf": False,
|
||||
"require": ["iss", "iat", "exp"],
|
||||
},
|
||||
)
|
||||
except jwt.InvalidSignatureError:
|
||||
return BadSignature()
|
||||
except (jwt.InvalidTokenError, ValueError, TypeError):
|
||||
return MalformedPayload()
|
||||
try:
|
||||
return _EnvelopeClaims.model_validate(payload)
|
||||
except ValidationError:
|
||||
return MalformedPayload()
|
||||
|
||||
|
||||
def _encrypt_grant_blob(plaintext: str, encryption_key: SecretStr) -> str:
|
||||
ciphertext = bytes(encrypt_value(value=plaintext, signing_key=encryption_key.get_secret_value()))
|
||||
return base64.urlsafe_b64encode(ciphertext).decode("ascii")
|
||||
|
||||
|
||||
def _decrypt_grant(
|
||||
blob: str,
|
||||
encryption_key: SecretStr,
|
||||
) -> UpstreamTokenGrant | DecryptFailed | MalformedPayload:
|
||||
from nacl.exceptions import CryptoError
|
||||
|
||||
try:
|
||||
plaintext = decrypt_value(
|
||||
value=base64.urlsafe_b64decode(blob),
|
||||
signing_key=encryption_key.get_secret_value(),
|
||||
)
|
||||
except (CryptoError, ValueError):
|
||||
return DecryptFailed()
|
||||
try:
|
||||
return UpstreamTokenGrant.model_validate_json(plaintext)
|
||||
except ValidationError:
|
||||
return MalformedPayload()
|
||||
|
|
@ -69,6 +69,17 @@ class OAuthTokenStore(Protocol):
|
|||
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ...
|
||||
|
||||
|
||||
class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol):
|
||||
"""An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped.
|
||||
|
||||
The write side calls ``invalidate`` after a (re)authorization or revocation changes the
|
||||
credential row, so reads stop serving the replaced token immediately instead of until its
|
||||
cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this.
|
||||
"""
|
||||
|
||||
async def invalidate(self, user_id: str, server_id: str) -> None: ...
|
||||
|
||||
|
||||
class TokenRefresher(Protocol):
|
||||
"""Mints a fresh token from an expired one and persists it, returning the new token.
|
||||
|
||||
|
|
|
|||
|
|
@ -24,8 +24,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_toke
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
CachedOAuthTokenStore,
|
||||
InvalidatableOAuthTokenStore,
|
||||
OAuthToken,
|
||||
OAuthTokenStore,
|
||||
RefreshCoordinator,
|
||||
RefreshingTokenStore,
|
||||
TokenCacheBackend,
|
||||
|
|
@ -51,7 +51,7 @@ if TYPE_CHECKING:
|
|||
_DEFAULT_TTL_SECONDS = 300.0
|
||||
|
||||
ServerLookup = Callable[[str], "MCPServer | None"]
|
||||
StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]]
|
||||
StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]
|
||||
|
||||
|
||||
async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
|
||||
|
|
@ -185,7 +185,7 @@ class LazyPerUserOAuthTokenStore:
|
|||
self._server_lookup = server_lookup
|
||||
self._store_builder = store_builder
|
||||
self._redis_available = redis_available
|
||||
self._store: OAuthTokenStore | None = None
|
||||
self._store: InvalidatableOAuthTokenStore | None = None
|
||||
self._uses_redis = False
|
||||
self._fetch_lock = asyncio.Condition()
|
||||
self._local_fetches = 0
|
||||
|
|
@ -203,7 +203,26 @@ class LazyPerUserOAuthTokenStore:
|
|||
if not uses_redis:
|
||||
await self._finish_local_fetch()
|
||||
|
||||
async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]:
|
||||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
"""Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row
|
||||
changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis)
|
||||
cache entry written by another worker is dropped too; the in-process case is then a no-op
|
||||
on an empty cache.
|
||||
"""
|
||||
if self._uses_redis:
|
||||
store = self._store
|
||||
if store is not None:
|
||||
await store.invalidate(user_id, server_id)
|
||||
return
|
||||
|
||||
store, uses_redis = await self._store_for_fetch()
|
||||
try:
|
||||
await store.invalidate(user_id, server_id)
|
||||
finally:
|
||||
if not uses_redis:
|
||||
await self._finish_local_fetch()
|
||||
|
||||
async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]:
|
||||
async with self._fetch_lock:
|
||||
while (
|
||||
self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0
|
||||
|
|
|
|||
|
|
@ -7,10 +7,11 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin
|
|||
an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
|
||||
at runtime instead of returning `None`.
|
||||
|
||||
`none` and `api_key` (shared-key source) are live, as is `authorization_code`, which reads the
|
||||
user's token from the injected `OAuthTokenStore`, and `token_exchange`, which swaps the caller's
|
||||
inbound token through the injected `TokenExchanger`. The remaining arms are `not_implemented` stubs
|
||||
that each land in a follow-up PR with their seam. Pure v2: no imports from v1.
|
||||
`none`, `api_key` (shared-key source), and `passthrough` (forwards the caller's own inbound token)
|
||||
are live, as is `authorization_code`, which reads the user's token from the injected
|
||||
`OAuthTokenStore`, and `token_exchange`, which swaps the caller's inbound token through the
|
||||
injected `TokenExchanger`. The remaining arms are `not_implemented` stubs that each land in a
|
||||
follow-up PR with their seam. Pure v2: no imports from v1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -97,7 +98,7 @@ class UpstreamCredentialProvider:
|
|||
case ApiKeyConfig() as config:
|
||||
return self._api_key(config)
|
||||
case PassthroughConfig():
|
||||
return _not_implemented(AuthSpecKind.passthrough)
|
||||
return self._passthrough(subject)
|
||||
case ClientCredentialsConfig():
|
||||
return _not_implemented(AuthSpecKind.client_credentials)
|
||||
case TokenExchangeConfig() as config:
|
||||
|
|
@ -118,6 +119,18 @@ class UpstreamCredentialProvider:
|
|||
"""
|
||||
return await self._authz_token(subject, server) is not None
|
||||
|
||||
def _passthrough(self, subject: Subject) -> Result[httpx.Auth, CredError]:
|
||||
"""Forward the caller's own upstream credential verbatim; the gateway mints nothing.
|
||||
|
||||
The inbound token is the caller's already-disambiguated ``Authorization`` (never the LiteLLM
|
||||
admission credential; the edge adapter drops that before building the ``Subject``). When it is
|
||||
absent the request is sent unauthenticated so the upstream's own 401 surfaces, rather than the
|
||||
gateway challenging on the upstream's behalf.
|
||||
"""
|
||||
if subject.inbound_token is None:
|
||||
return Ok(NoOpAuth())
|
||||
return Ok(StaticHeaderAuth(subject.inbound_token.get_secret_value(), header_name="Authorization"))
|
||||
|
||||
def _api_key(self, config: ApiKeyConfig) -> Result[httpx.Auth, CredError]:
|
||||
match config.key_source:
|
||||
case SharedKey() as source:
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
|||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
|
||||
|
||||
|
||||
class AuthSpecKind(str, Enum):
|
||||
|
|
@ -215,7 +216,7 @@ class TokenExchangeConfig(BaseModel):
|
|||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal[AuthSpecKind.token_exchange] = AuthSpecKind.token_exchange
|
||||
profile: Literal["rfc8693", "entra_obo"] = "rfc8693"
|
||||
subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token"
|
||||
subject_token_type: str = DEFAULT_SUBJECT_TOKEN_TYPE
|
||||
token_exchange_endpoint: str | None = None
|
||||
audience: str | None = None
|
||||
client_id: str | None = None
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue