mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into feature/ovalix-extended-guardrail
This commit is contained in:
commit
842525f2ba
508 changed files with 37534 additions and 10221 deletions
|
|
@ -2731,7 +2731,7 @@ jobs:
|
|||
- ~/.cache/uv
|
||||
- restore_cache:
|
||||
keys:
|
||||
- ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
- ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }}
|
||||
- run:
|
||||
name: Install Node dependencies and Playwright
|
||||
# The cimg/python:3.12-browsers image already ships the Chromium system
|
||||
|
|
@ -2742,11 +2742,14 @@ jobs:
|
|||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npm ci
|
||||
cd ../../tests/e2e/ui
|
||||
npm ci
|
||||
npx playwright install chromium
|
||||
- save_cache:
|
||||
key: ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
key: ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }}
|
||||
paths:
|
||||
- ui/litellm-dashboard/node_modules
|
||||
- tests/e2e/ui/node_modules
|
||||
- ~/.cache/ms-playwright
|
||||
- run:
|
||||
name: Build UI from source
|
||||
|
|
@ -2777,10 +2780,10 @@ jobs:
|
|||
name: Seed database
|
||||
command: |
|
||||
PGPASSWORD=e2epassword psql -h localhost -p 5432 -U e2euser -d litellm_e2e \
|
||||
-f ui/litellm-dashboard/e2e_tests/fixtures/seed.sql
|
||||
-f tests/e2e/ui/fixtures/seed.sql
|
||||
- run:
|
||||
name: Start mock LLM server
|
||||
command: uv run --no-sync python ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Start LiteLLM proxy
|
||||
|
|
@ -2798,7 +2801,7 @@ jobs:
|
|||
command: |
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--config tests/e2e/ui/fixtures/config.yml \
|
||||
--port 4000
|
||||
background: true
|
||||
- run:
|
||||
|
|
@ -2819,15 +2822,15 @@ jobs:
|
|||
# Forward LITELLM_LICENSE so license.spec.ts can detect that the
|
||||
# proxy was launched with a license and assert premium_user=true.
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
cd tests/e2e/ui
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
npx playwright test --config playwright.config.ts
|
||||
no_output_timeout: 10m
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/test-results
|
||||
path: tests/e2e/ui/test-results
|
||||
destination: e2e-test-results
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/playwright-report
|
||||
path: tests/e2e/ui/playwright-report
|
||||
destination: e2e-playwright-report
|
||||
|
||||
e2e_ui_testing_server_root_path:
|
||||
|
|
@ -2870,17 +2873,20 @@ jobs:
|
|||
- ~/.cache/uv
|
||||
- restore_cache:
|
||||
keys:
|
||||
- ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
- ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }}
|
||||
- run:
|
||||
name: Install Node dependencies and Playwright
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npm ci
|
||||
cd ../../tests/e2e/ui
|
||||
npm ci
|
||||
npx playwright install chromium
|
||||
- save_cache:
|
||||
key: ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
key: ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }}
|
||||
paths:
|
||||
- ui/litellm-dashboard/node_modules
|
||||
- tests/e2e/ui/node_modules
|
||||
- ~/.cache/ms-playwright
|
||||
- run:
|
||||
name: Build UI from source
|
||||
|
|
@ -2902,10 +2908,10 @@ jobs:
|
|||
name: Seed database
|
||||
command: |
|
||||
PGPASSWORD=e2epassword psql -h localhost -p 5432 -U e2euser -d litellm_e2e \
|
||||
-f ui/litellm-dashboard/e2e_tests/fixtures/seed.sql
|
||||
-f tests/e2e/ui/fixtures/seed.sql
|
||||
- run:
|
||||
name: Start mock LLM server
|
||||
command: uv run --no-sync python ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Start LiteLLM proxy under a server root path
|
||||
|
|
@ -2918,7 +2924,7 @@ jobs:
|
|||
command: |
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--config tests/e2e/ui/fixtures/config.yml \
|
||||
--port 4000
|
||||
background: true
|
||||
- run:
|
||||
|
|
@ -2937,15 +2943,15 @@ jobs:
|
|||
- run:
|
||||
name: Run migration smoke under SERVER_ROOT_PATH
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
cd tests/e2e/ui
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
npx playwright test --config e2e_tests/migration.serverRootPath.config.ts
|
||||
npx playwright test --config migration.serverRootPath.config.ts
|
||||
no_output_timeout: 10m
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/test-results
|
||||
path: tests/e2e/ui/test-results
|
||||
destination: e2e-server-root-path-test-results
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/playwright-report
|
||||
path: tests/e2e/ui/playwright-report
|
||||
destination: e2e-server-root-path-playwright-report
|
||||
|
||||
build_docker_database_image:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ has_backend=false
|
|||
while IFS= read -r file || [ -n "$file" ]; do
|
||||
[ -n "$file" ] || continue
|
||||
case "$file" in
|
||||
ui/*) has_client=true ;;
|
||||
ui/* | tests/e2e/ui/*) has_client=true ;;
|
||||
docs/* | *.md | *.mdx) : ;;
|
||||
*) has_backend=true ;;
|
||||
esac
|
||||
|
|
|
|||
2
.github/ISSUE_TEMPLATE/bug_report.yml
vendored
2
.github/ISSUE_TEMPLATE/bug_report.yml
vendored
|
|
@ -30,7 +30,7 @@ body:
|
|||
id: steps-to-reproduce
|
||||
attributes:
|
||||
label: Steps to Reproduce
|
||||
description: Please provide detailed steps to reproduce this bug(A curl/python code to reproduce the bug)
|
||||
description: Please provide a numbered list of the exact steps to reproduce this bug (include a curl/python snippet to reproduce it). Number each step (1., 2., 3., ...) in the order you performed them.
|
||||
placeholder: |
|
||||
1. config.yaml file/ .env file/ etc.
|
||||
2. Run the following code...
|
||||
|
|
|
|||
15
.github/pull_request_template.md
vendored
15
.github/pull_request_template.md
vendored
|
|
@ -1,3 +1,18 @@
|
|||
## TLDR
|
||||
|
||||
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max
|
||||
This section must be extremely human parsable, comprehensible, and readable: its target audience is humans, not AI agents -->
|
||||
|
||||
Problem this solves:
|
||||
|
||||
- <blah>
|
||||
- ...
|
||||
|
||||
How it solves it:
|
||||
|
||||
- <blah>
|
||||
- ...
|
||||
|
||||
## Relevant issues
|
||||
|
||||
<!-- e.g., "Fixes #000" -->
|
||||
|
|
|
|||
20
.github/workflows/image-scan.yml
vendored
20
.github/workflows/image-scan.yml
vendored
|
|
@ -9,6 +9,7 @@ on:
|
|||
- "litellm_**"
|
||||
paths:
|
||||
- docker/Dockerfile.non_root
|
||||
- tests/proxy_migration_tests/test_offline_image_migration.py
|
||||
- uv.lock
|
||||
- ui/litellm-dashboard/package-lock.json
|
||||
- .github/workflows/image-scan.yml
|
||||
|
|
@ -51,6 +52,23 @@ jobs:
|
|||
- name: Build runtime image
|
||||
run: docker build -f docker/Dockerfile.non_root -t litellm-image-scan:${{ github.sha }} .
|
||||
|
||||
# The prisma bake must migrate a fresh DB with no egress as an arbitrary
|
||||
# non-root uid (OpenShift restricted-v2 / air-gapped / readOnlyRootFilesystem).
|
||||
# `docker run` as the default uid with network hides a broken bake because
|
||||
# the migration entrypoint exits 0 even when it applied nothing; asserting
|
||||
# the schema was created is what catches it.
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Verify offline migration as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-image-scan:${{ github.sha }}
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
|
||||
|
||||
# Scans the whole shipped artifact: OS/apk plus every language package
|
||||
# baked into the image, including ones no lockfile declares (e.g. prisma's
|
||||
# vendored node engine) that osv-scan cannot see. osv-scan stays the fast
|
||||
|
|
@ -58,6 +76,8 @@ jobs:
|
|||
# free OSS, run as a pinned, checksum-verified binary; no GitHub Action
|
||||
# dependency and no vendor SaaS callout.
|
||||
- name: Scan image for fixable HIGH/CRITICAL CVEs
|
||||
env:
|
||||
GRYPE_MATCH_PYTHON_USING_CPES: "true"
|
||||
run: |
|
||||
"$RUNNER_TEMP/grype" litellm-image-scan:${{ github.sha }} \
|
||||
--only-fixed \
|
||||
|
|
|
|||
3
.github/workflows/test-code-quality.yml
vendored
3
.github/workflows/test-code-quality.yml
vendored
|
|
@ -115,6 +115,9 @@ jobs:
|
|||
- name: check_fastuuid_usage
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py
|
||||
|
||||
- name: check_e2e_no_raw_requests
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py
|
||||
|
||||
- name: memory_test
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
|
||||
|
||||
|
|
|
|||
57
.github/workflows/test-litellm-ui-unit.yml
vendored
Normal file
57
.github/workflows/test-litellm-ui-unit.yml
vendored
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
name: UI Unit Tests
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- litellm_internal_staging
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
ui-unit-tests:
|
||||
runs-on: ubuntu-latest-16-cores
|
||||
timeout-minutes: 20
|
||||
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: Setup Node.js
|
||||
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
|
||||
run: npm ci
|
||||
|
||||
- name: Run UI unit tests (Vitest)
|
||||
env:
|
||||
CI: "true"
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
if [ -n "$BASE_SHA" ]; then
|
||||
echo "Pull request: running only tests related to changes since $BASE_SHA"
|
||||
npm run test -- --run --changed "$BASE_SHA" --passWithNoTests \
|
||||
--pool forks --poolOptions.forks.maxForks=14
|
||||
else
|
||||
echo "Push to $GITHUB_REF_NAME: running the full suite"
|
||||
npm run test -- --run --pool forks --poolOptions.forks.maxForks=14
|
||||
fi
|
||||
|
|
@ -46,6 +46,7 @@ jobs:
|
|||
tests/test_litellm/proxy/rag_endpoints
|
||||
tests/test_litellm/proxy/realtime_endpoints
|
||||
tests/test_litellm/proxy/ui_crud_endpoints
|
||||
tests/test_litellm/proxy/config_resolvers
|
||||
tests/test_litellm/proxy/utils
|
||||
workers: 2
|
||||
reruns: 2
|
||||
|
|
|
|||
10
.github/workflows/test_server_root_path.yml
vendored
10
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -106,8 +106,8 @@ jobs:
|
|||
with:
|
||||
node-version: "20"
|
||||
|
||||
- name: Install UI deps and Chromium
|
||||
working-directory: ui/litellm-dashboard
|
||||
- name: Install e2e deps and Chromium
|
||||
working-directory: tests/e2e/ui
|
||||
run: |
|
||||
retry() {
|
||||
local attempt=1
|
||||
|
|
@ -131,17 +131,17 @@ jobs:
|
|||
retry npx playwright install --with-deps chromium
|
||||
|
||||
- name: Run SERVER_ROOT_PATH redirect e2e
|
||||
working-directory: ui/litellm-dashboard
|
||||
working-directory: tests/e2e/ui
|
||||
env:
|
||||
SERVER_ROOT_PATH: ${{ matrix.root_path }}
|
||||
run: npx playwright test --config=e2e_tests/serverRootPath.config.ts
|
||||
run: npx playwright test --config=serverRootPath.config.ts
|
||||
|
||||
- name: Upload Playwright artifacts on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2
|
||||
with:
|
||||
name: playwright-trace-${{ strategy.job-index }}
|
||||
path: ui/litellm-dashboard/test-results/
|
||||
path: tests/e2e/ui/test-results/
|
||||
retention-days: 7
|
||||
|
||||
- name: Cleanup
|
||||
|
|
|
|||
81
.github/workflows/weekly_load_anomaly.yml
vendored
Normal file
81
.github/workflows/weekly_load_anomaly.yml
vendored
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
name: "Weekly Load Anomaly Check"
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: "0 12 * * 6"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
weekly-load-anomaly:
|
||||
if: github.event_name != 'schedule' || github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 45
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:16.6
|
||||
env:
|
||||
POSTGRES_USER: llmproxy
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
POSTGRES_DB: litellm
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U llmproxy"
|
||||
--health-interval 5s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
env:
|
||||
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
LITELLM_MASTER_KEY: sk-weekly-anomaly-check
|
||||
ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }}
|
||||
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.AWS_BEARER_TOKEN_BEDROCK }}
|
||||
steps:
|
||||
- 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: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
|
||||
|
||||
- name: Generate Prisma client
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Start the proxy
|
||||
run: |
|
||||
nohup uv run --no-sync litellm --config tests/e2e/load/weekly_anomaly_config.yml --port 4000 > proxy.log 2>&1 &
|
||||
for _ in $(seq 1 90); do
|
||||
if curl -fs http://localhost:4000/health/liveliness > /dev/null; then
|
||||
exit 0
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
echo "proxy never became live"
|
||||
tail -n 100 proxy.log
|
||||
exit 1
|
||||
|
||||
- name: Run the weekly session anomaly test
|
||||
env:
|
||||
E2E_WEEKLY_ANOMALY: "1"
|
||||
run: |
|
||||
uv run --no-sync pytest tests/e2e/load/test_weekly_session_anomaly_e2e.py -v --tb=short -rA
|
||||
|
||||
- name: Show proxy log on failure
|
||||
if: failure()
|
||||
run: tail -n 300 proxy.log
|
||||
|
|
@ -18,6 +18,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/team/",
|
||||
"/v2/team/",
|
||||
"/organization/",
|
||||
"/v2/organization/",
|
||||
"/customer/",
|
||||
"/end_user/",
|
||||
"/sso/",
|
||||
|
|
|
|||
|
|
@ -54,7 +54,6 @@ ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
|||
UV_LINK_MODE=copy \
|
||||
PATH="/app/.venv/bin:${PATH}" \
|
||||
LITELLM_NON_ROOT=true \
|
||||
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
|
||||
XDG_CACHE_HOME=/app/.cache
|
||||
|
||||
# Copy dependency metadata first for layer caching
|
||||
|
|
@ -106,7 +105,9 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--python python3; \
|
||||
fi
|
||||
|
||||
RUN prisma generate --schema=./schema.prisma
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
prisma generate --schema=./schema.prisma
|
||||
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
|
||||
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
|
|
@ -127,8 +128,6 @@ RUN for i in 1 2 3; do \
|
|||
# the rest of the builder's /app is source and build metadata that must not
|
||||
# ship (manifest-scanning tools attribute everything in it to this image).
|
||||
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
|
||||
# Prisma caches live under /app/.cache here (XDG_CACHE_HOME /
|
||||
# PRISMA_BINARY_CACHE_DIR) so the runtime prisma generate finds them.
|
||||
COPY --from=builder /app/.venv /app/.venv
|
||||
COPY --from=builder /app/docker /app/docker
|
||||
COPY --from=builder /app/schema.prisma /app/schema.prisma
|
||||
|
|
@ -138,21 +137,35 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
|
|||
# enterprise.enterprise_hooks from it)
|
||||
COPY --from=builder /app/enterprise /app/enterprise
|
||||
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
|
||||
COPY --from=builder /app/.cache /app/.cache
|
||||
# Prisma CLI + engines are baked under /opt/prisma, a fixed path every runtime
|
||||
# uid can read and that no cache volume mount shadows (unlike /app/.cache or
|
||||
# $HOME/.cache under readOnlyRootFilesystem + emptyDir or arbitrary-uid setups).
|
||||
# PRISMA_CLI_QUERY_ENGINE_TYPE=binary makes the CLI use the baked binary query
|
||||
# engine directly, so `prisma migrate deploy` on a fresh database needs no npm
|
||||
# and no network access; without it the CLI looks for the library engine, which
|
||||
# prisma stopped baking, and falls back to a download that fails offline or as a
|
||||
# non-writable uid (#33650, #24554).
|
||||
COPY --from=builder /opt/prisma /opt/prisma
|
||||
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
|
||||
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets
|
||||
|
||||
# XDG_CACHE_HOME is intentionally left unset so it falls back to $HOME/.cache
|
||||
# (/app/.cache, writable by the runtime uid). The prisma bake at the read-only
|
||||
# /opt/prisma is anchored by PRISMA_BINARY_CACHE_DIR / PRISMA_CLI_PATH, so
|
||||
# nothing needs XDG to point there; pointing it at the read-only bake would
|
||||
# deny any XDG-aware library that writes a cache at runtime.
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
|
||||
PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
PRISMA_CLI_PATH=/opt/prisma/binaries/node_modules/.bin/prisma \
|
||||
PRISMA_CLI_QUERY_ENGINE_TYPE=binary \
|
||||
HOME=/app \
|
||||
LITELLM_NON_ROOT=true \
|
||||
XDG_CACHE_HOME=/app/.cache \
|
||||
PRISMA_SKIP_POSTINSTALL_GENERATE=1 \
|
||||
PRISMA_HIDE_UPDATE_MESSAGE=1 \
|
||||
PRISMA_ENGINES_CHECKSUM_IGNORE_MISSING=1 \
|
||||
PRISMA_OFFLINE_MODE=true
|
||||
|
||||
RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
|
||||
RUN mkdir -p /nonexistent /app/.cache /var/lib/litellm/assets /var/lib/litellm/ui && \
|
||||
chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent && \
|
||||
PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \
|
||||
chown -R nobody:nogroup "$PRISMA_PATH" && \
|
||||
|
|
@ -165,12 +178,14 @@ RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
|
|||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g=u "$LITELLM_PROXY_EXTRAS_PATH" || true && \
|
||||
chmod -R g+w "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \
|
||||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w "$LITELLM_PROXY_EXTRAS_PATH" || true && \
|
||||
chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets /app/.cache
|
||||
chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
test -x /opt/prisma/binaries/node_modules/.bin/prisma && \
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \
|
||||
ls /opt/prisma/binaries/node_modules/@prisma/engines/query-engine-* >/dev/null 2>&1
|
||||
|
||||
USER 65534
|
||||
|
||||
RUN prisma generate --schema=./schema.prisma
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
ENTRYPOINT ["/app/docker/prod_entrypoint.sh"]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,9 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_SSOIdentityAssertion" (
|
||||
"user_id" TEXT NOT NULL,
|
||||
"assertion_b64" TEXT NOT NULL,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_SSOIdentityAssertion_pkey" PRIMARY KEY ("user_id")
|
||||
);
|
||||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.79"
|
||||
version = "0.4.80"
|
||||
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.79"
|
||||
version = "0.4.80"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -50,3 +50,15 @@ pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool {
|
|||
.iter()
|
||||
.any(|(key, _)| key.eq_ignore_ascii_case(name))
|
||||
}
|
||||
|
||||
pub(super) fn has_bearer_auth(headers: &[(String, String)]) -> bool {
|
||||
headers.iter().any(|(name, value)| {
|
||||
if !name.eq_ignore_ascii_case("authorization") {
|
||||
return false;
|
||||
}
|
||||
let value = value.trim();
|
||||
value.len() > 7
|
||||
&& value[..7].eq_ignore_ascii_case("bearer ")
|
||||
&& !value[7..].trim().is_empty()
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use litellm_core::CoreResult;
|
|||
use litellm_core::messages::transformation::MessagesAuthStrategy;
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use super::common_utils::{has_header, messages_provider_config, string_headers};
|
||||
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
|
||||
use super::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) fn prepare_messages_call(
|
||||
|
|
@ -33,7 +33,9 @@ pub(super) fn prepare_messages_call(
|
|||
let mut headers = string_headers(request.extra_headers)?;
|
||||
|
||||
let auth_strategy = config.auth_strategy();
|
||||
if !has_header(&headers, auth_strategy.header_name()) {
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
use super::common_utils::{
|
||||
has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
};
|
||||
use super::{MessagesRequest, messages};
|
||||
|
||||
|
|
@ -85,6 +85,34 @@ fn has_header_is_case_insensitive() {
|
|||
assert!(!has_header(&headers, "authorization"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn has_bearer_auth_requires_a_nonempty_bearer_token() {
|
||||
assert!(has_bearer_auth(&[(
|
||||
"Authorization".to_string(),
|
||||
"Bearer tok".to_string()
|
||||
)]));
|
||||
assert!(has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"bearer tok".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"Bearer ".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
String::new()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"Basic abc".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"x-api-key".to_string(),
|
||||
"sk".to_string()
|
||||
)]));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_round_trip_builds_azure_request_and_passes_response_through() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
|
|
@ -252,6 +280,112 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
|||
assert!(!head.contains("rust-fallback-key"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let request = read_http_request(&mut socket).await;
|
||||
let response_body =
|
||||
r#"{"id":"msg_3","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let mut headers = Map::new();
|
||||
headers.insert(
|
||||
"Authorization".to_string(),
|
||||
Value::String("Bearer entra-token".to_string()),
|
||||
);
|
||||
|
||||
messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: None,
|
||||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
})
|
||||
.await
|
||||
.expect("entra id request succeeds without api key");
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let head = request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("has body")
|
||||
.0
|
||||
.to_ascii_lowercase();
|
||||
assert!(head.contains("authorization: bearer entra-token"), "{head}");
|
||||
assert!(!head.contains("x-api-key"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_requires_auth_when_no_key_and_no_header() {
|
||||
let err = messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: None,
|
||||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
})
|
||||
.await
|
||||
.expect_err("missing auth errors");
|
||||
|
||||
assert!(matches!(err, CoreError::Auth(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let request = read_http_request(&mut socket).await;
|
||||
let response_body =
|
||||
r#"{"id":"msg_4","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let mut headers = Map::new();
|
||||
headers.insert(
|
||||
"Authorization".to_string(),
|
||||
Value::String("Bearer ".to_string()),
|
||||
);
|
||||
|
||||
messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("sk-azure"),
|
||||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
})
|
||||
.await
|
||||
.expect("falls back to api key");
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let head = request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("has body")
|
||||
.0
|
||||
.to_ascii_lowercase();
|
||||
assert!(head.contains("x-api-key: sk-azure"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_maps_provider_error_status_to_http_error() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
|
|
|
|||
|
|
@ -35,6 +35,10 @@ pub trait AnthropicMessagesProviderConfig: Sync {
|
|||
MessagesAuthStrategy::Header("x-api-key")
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("anthropic-version", "2023-06-01"),
|
||||
|
|
|
|||
|
|
@ -163,6 +163,10 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
|
|||
self.anthropic.auth_strategy()
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
self.anthropic.default_headers()
|
||||
}
|
||||
|
|
@ -294,6 +298,11 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_bearer_auth_for_entra_id() {
|
||||
assert!(AZURE_ANTHROPIC_MESSAGES_CONFIG.accepts_bearer_auth());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_headers_match_python() {
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -427,7 +427,9 @@ default_team_settings: Optional[List] = None
|
|||
max_user_budget: Optional[float] = None
|
||||
default_max_internal_user_budget: Optional[float] = None
|
||||
max_internal_user_budget: Optional[float] = None
|
||||
max_ui_session_budget: Optional[float] = 0.25 # $0.25 USD budgets for UI Chat sessions
|
||||
max_ui_session_budget: Optional[float] = (
|
||||
1.0 # USD budget for each dashboard login session (playground, test connection)
|
||||
)
|
||||
internal_user_budget_duration: Optional[str] = None
|
||||
tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None
|
||||
max_end_user_budget: Optional[float] = None
|
||||
|
|
|
|||
|
|
@ -264,6 +264,9 @@ MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT",
|
|||
# Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth
|
||||
LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000))
|
||||
TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60))
|
||||
GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS = int(
|
||||
os.getenv("GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS", 24 * 60 * 60)
|
||||
)
|
||||
# Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger.
|
||||
# Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire.
|
||||
MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8)))
|
||||
|
|
@ -1469,6 +1472,7 @@ _batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower()
|
|||
PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true"
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605))
|
||||
PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30)
|
||||
|
||||
# APScheduler Configuration - MEMORY LEAK FIX
|
||||
# These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions
|
||||
|
|
@ -1524,6 +1528,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
|
|||
# test_general_settings_ui_fields_are_db_overridable enforces that pairing.
|
||||
"enable_anthropic_prompt_caching",
|
||||
"anthropic_prompt_caching_ttl",
|
||||
"max_ui_session_budget",
|
||||
]
|
||||
SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"]
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
from datetime import datetime
|
||||
|
|
@ -46,7 +47,10 @@ if TYPE_CHECKING:
|
|||
dc = DualCache()
|
||||
|
||||
|
||||
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
|
||||
from litellm.constants import (
|
||||
GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
)
|
||||
from litellm.exceptions import (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
|
|
@ -113,6 +117,7 @@ class CustomGuardrail(CustomLogger):
|
|||
on_sensitive_data: Optional[str] = None,
|
||||
sensitive_data_route_to_model: Optional[str] = None,
|
||||
sticky_session_routing: bool = True,
|
||||
only_scan_new_messages: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -145,6 +150,7 @@ class CustomGuardrail(CustomLogger):
|
|||
self.on_sensitive_data: Optional[str] = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model
|
||||
self.sticky_session_routing: bool = sticky_session_routing
|
||||
self.only_scan_new_messages: bool = only_scan_new_messages
|
||||
|
||||
if supported_event_hooks:
|
||||
## validate event_hook is in supported_event_hooks
|
||||
|
|
@ -269,6 +275,100 @@ class CustomGuardrail(CustomLogger):
|
|||
"""Extract session_id from request data."""
|
||||
return get_session_id_from_request_data(request_data)
|
||||
|
||||
@staticmethod
|
||||
def _scanned_text_hash(text: str) -> str:
|
||||
"""Stable content hash for a single scannable text segment.
|
||||
|
||||
Hashing the exact text the provider would receive means an edited earlier
|
||||
segment produces a different hash and gets re-scanned, while an unchanged
|
||||
segment repeated on a later turn is skipped.
|
||||
"""
|
||||
return hashlib.sha256(text.encode("utf-8")).hexdigest()
|
||||
|
||||
def _scanned_texts_cache_key(self, session_id: str) -> str:
|
||||
return f"guardrail_scanned_texts:{self.guardrail_name}:{session_id}"
|
||||
|
||||
async def filter_new_texts_for_session(
|
||||
self,
|
||||
texts: list[str] | None,
|
||||
request_data: dict[str, object],
|
||||
cache: DualCache,
|
||||
) -> list[str] | None:
|
||||
"""Return only the text segments not already scanned earlier in this session.
|
||||
|
||||
Returns ``None`` when incremental scanning is inactive (feature off, no
|
||||
session id, masking enabled, or the cache read failed). ``None`` signals
|
||||
the caller to fall back to a full scan; a returned list (possibly empty)
|
||||
signals the caller to scan only that subset and skip masking write-back.
|
||||
"""
|
||||
if not self.only_scan_new_messages or not texts:
|
||||
return None
|
||||
|
||||
if self.mask_request_content or self.mask_response_content:
|
||||
verbose_logger.warning(
|
||||
"Guardrail %s: only_scan_new_messages is not supported with masking; scanning full context.",
|
||||
self.guardrail_name,
|
||||
)
|
||||
return None
|
||||
|
||||
session_id = get_session_id_from_request_data(request_data)
|
||||
if not session_id:
|
||||
verbose_logger.debug(
|
||||
"Guardrail %s: only_scan_new_messages enabled but request has no session id; scanning full context.",
|
||||
self.guardrail_name,
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
cached: object = await cache.async_get_cache(key=self._scanned_texts_cache_key(session_id))
|
||||
except Exception as e: # noqa: BLE001 # cache is best-effort; any failure must fall back to a full scan
|
||||
verbose_logger.warning(
|
||||
"Guardrail %s: failed to read scanned-message cache (%s); scanning full context.",
|
||||
self.guardrail_name,
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
||||
seen: set[str] = {str(h) for h in cached} if isinstance(cached, list) else set()
|
||||
return [text for text in texts if self._scanned_text_hash(text) not in seen]
|
||||
|
||||
async def mark_texts_scanned(
|
||||
self,
|
||||
texts: list[str] | None,
|
||||
request_data: dict[str, object],
|
||||
cache: DualCache,
|
||||
) -> None:
|
||||
"""Record the hashes of all text segments present on a successful (non-blocked) scan.
|
||||
|
||||
Called only after the guardrail allows the request, so a blocked segment is
|
||||
never marked scanned and will be re-checked if the client retries.
|
||||
"""
|
||||
if not self.only_scan_new_messages or not texts:
|
||||
return
|
||||
if self.mask_request_content or self.mask_response_content:
|
||||
return
|
||||
session_id = get_session_id_from_request_data(request_data)
|
||||
if not session_id:
|
||||
return
|
||||
|
||||
cache_key = self._scanned_texts_cache_key(session_id)
|
||||
current_hashes = [self._scanned_text_hash(text) for text in texts]
|
||||
try:
|
||||
existing: object = await cache.async_get_cache(key=cache_key)
|
||||
existing_hashes: list[str] = [str(h) for h in existing] if isinstance(existing, list) else []
|
||||
merged: list[str] = list(dict.fromkeys(existing_hashes + current_hashes))
|
||||
await cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=merged,
|
||||
ttl=GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # cache is best-effort; any failure must not block the request
|
||||
verbose_logger.warning(
|
||||
"Guardrail %s: failed to persist scanned-message cache (%s); next call will re-scan.",
|
||||
self.guardrail_name,
|
||||
e,
|
||||
)
|
||||
|
||||
def should_route_on_sensitive_data(self) -> bool:
|
||||
"""
|
||||
Returns True if this guardrail is configured to route requests
|
||||
|
|
|
|||
|
|
@ -7,11 +7,24 @@ duration_in_seconds is used in diff parts of the code base, example
|
|||
"""
|
||||
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone, tzinfo
|
||||
from typing import Optional, Tuple
|
||||
import time as time_module
|
||||
from datetime import datetime, time, timedelta, timezone, tzinfo
|
||||
from typing import Final, Optional, Tuple
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
_BUDGET_DURATION_WORD_ALIASES: Final[dict[str, str]] = {
|
||||
"hourly": "1h",
|
||||
"daily": "24h",
|
||||
"weekly": "7d",
|
||||
"monthly": "30d",
|
||||
}
|
||||
|
||||
|
||||
def _normalize_duration(duration: str) -> str:
|
||||
return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration)
|
||||
|
||||
|
||||
def _extract_from_regex(duration: str) -> Tuple[int, str]:
|
||||
match = re.match(r"(\d+)(mo|[smhdw]?)", duration)
|
||||
|
|
@ -48,7 +61,7 @@ def duration_in_seconds(duration: str) -> int:
|
|||
|
||||
Returns time in seconds till when budget needs to be reset
|
||||
"""
|
||||
value, unit = _extract_from_regex(duration=duration)
|
||||
value, unit = _extract_from_regex(duration=_normalize_duration(duration))
|
||||
|
||||
if unit == "s":
|
||||
return value
|
||||
|
|
@ -61,7 +74,7 @@ def duration_in_seconds(duration: str) -> int:
|
|||
elif unit == "w":
|
||||
return value * 604800
|
||||
elif unit == "mo":
|
||||
now = time.time()
|
||||
now = time_module.time()
|
||||
current_time = datetime.fromtimestamp(now)
|
||||
|
||||
# Calculate target month and year, handling overflow past December
|
||||
|
|
@ -94,12 +107,17 @@ def duration_in_seconds(duration: str) -> int:
|
|||
raise ValueError(f"Unsupported duration unit, passed duration: {duration}")
|
||||
|
||||
|
||||
def get_next_standardized_reset_time(duration: str, current_time: datetime, timezone_str: str = "UTC") -> datetime:
|
||||
def get_next_standardized_reset_time(
|
||||
duration: str,
|
||||
current_time: datetime,
|
||||
timezone_str: str = "UTC",
|
||||
reset_time_of_day: time = time(0, 0),
|
||||
) -> datetime:
|
||||
"""
|
||||
Get the next standardized reset time based on the duration.
|
||||
|
||||
All durations will reset at predictable intervals, aligned from the current time:
|
||||
- Nd: If N=1, reset at next midnight; if N>1, reset every N days from now
|
||||
- Nd: If N=1, reset at the next `reset_time_of_day`; if N>1, reset every N days from now
|
||||
- Nh: Every N hours, aligned to hour boundaries (e.g., 1:00, 2:00)
|
||||
- Nm: Every N minutes, aligned to minute boundaries (e.g., 1:05, 1:10)
|
||||
- Ns: Every N seconds, aligned to second boundaries
|
||||
|
|
@ -108,17 +126,24 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
|
|||
- duration: Duration string (e.g. "30s", "30m", "30h", "30d")
|
||||
- current_time: Current datetime
|
||||
- timezone_str: Timezone string (e.g. "UTC", "US/Eastern", "Asia/Kolkata")
|
||||
- reset_time_of_day: Wall-clock time the reset lands on for day/week/month
|
||||
durations (defaults to midnight). Ignored for sub-day durations, where a
|
||||
time-of-day is meaningless.
|
||||
|
||||
Returns:
|
||||
- Next reset time at a standardized interval in the specified timezone
|
||||
"""
|
||||
# Set up timezone and normalize current time
|
||||
current_time, tz = _setup_timezone(current_time, timezone_str)
|
||||
current_time, _ = _setup_timezone(current_time, timezone_str)
|
||||
|
||||
# Parse duration
|
||||
value, unit = _parse_duration(duration)
|
||||
value, unit = _parse_duration(_normalize_duration(duration))
|
||||
if value is None:
|
||||
# Fall back to default if format is invalid
|
||||
verbose_logger.warning(
|
||||
"Unrecognized budget_duration %r; falling back to a next-midnight reset. "
|
||||
"Use the <int><unit> format (e.g. '1h', '7d', '30d', '1mo').",
|
||||
duration,
|
||||
)
|
||||
return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=1)
|
||||
|
||||
# Midnight of the current day in the specified timezone
|
||||
|
|
@ -126,9 +151,9 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
|
|||
|
||||
# Handle different time units
|
||||
if unit == "d":
|
||||
return _handle_day_reset(current_time, base_midnight, value, tz)
|
||||
return _handle_day_reset(current_time, base_midnight, value, reset_time_of_day)
|
||||
elif unit == "w":
|
||||
return _handle_day_reset(current_time, base_midnight, value * 7, tz)
|
||||
return _handle_day_reset(current_time, base_midnight, value * 7, reset_time_of_day)
|
||||
elif unit == "h":
|
||||
return _handle_hour_reset(current_time, base_midnight, value)
|
||||
elif unit == "m":
|
||||
|
|
@ -136,7 +161,7 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
|
|||
elif unit == "s":
|
||||
return _handle_second_reset(current_time, base_midnight, value)
|
||||
elif unit == "mo":
|
||||
return _handle_month_reset(current_time, base_midnight, value)
|
||||
return _handle_month_reset(current_time, base_midnight, value, reset_time_of_day)
|
||||
else:
|
||||
# Unrecognized unit, default to next midnight
|
||||
return base_midnight + timedelta(days=1)
|
||||
|
|
@ -175,46 +200,58 @@ def _parse_duration(duration: str) -> Tuple[Optional[int], Optional[str]]:
|
|||
return int(value), unit
|
||||
|
||||
|
||||
def _handle_day_reset(current_time: datetime, base_midnight: datetime, value: int, tz: tzinfo) -> datetime:
|
||||
def _apply_time_of_day(dt: datetime, reset_time_of_day: time) -> datetime:
|
||||
"""Set the wall-clock time of `dt` to `reset_time_of_day`, keeping its date and tzinfo."""
|
||||
return dt.replace(
|
||||
hour=reset_time_of_day.hour,
|
||||
minute=reset_time_of_day.minute,
|
||||
second=reset_time_of_day.second,
|
||||
microsecond=reset_time_of_day.microsecond,
|
||||
)
|
||||
|
||||
|
||||
def _next_occurrence(
|
||||
boundary_midnight: datetime,
|
||||
reset_time_of_day: time,
|
||||
current_time: datetime,
|
||||
period: timedelta,
|
||||
) -> datetime:
|
||||
"""Place the reset at `reset_time_of_day` on the boundary day, rolling forward one
|
||||
`period` if that instant has already passed (or is exactly now)."""
|
||||
candidate = _apply_time_of_day(boundary_midnight, reset_time_of_day)
|
||||
if candidate <= current_time:
|
||||
return candidate + period
|
||||
return candidate
|
||||
|
||||
|
||||
def _first_of_next_month(first_of_month: datetime) -> datetime:
|
||||
"""Given the 1st of some month, return the 1st of the following month."""
|
||||
if first_of_month.month == 12:
|
||||
return first_of_month.replace(year=first_of_month.year + 1, month=1)
|
||||
return first_of_month.replace(month=first_of_month.month + 1)
|
||||
|
||||
|
||||
def _handle_day_reset(
|
||||
current_time: datetime,
|
||||
base_midnight: datetime,
|
||||
value: int,
|
||||
reset_time_of_day: time,
|
||||
) -> datetime:
|
||||
"""Handle day-based reset times."""
|
||||
# Handle zero value - immediate expiration
|
||||
if value == 0:
|
||||
return current_time
|
||||
|
||||
if value == 1: # Daily reset at midnight
|
||||
return base_midnight + timedelta(days=1)
|
||||
elif value == 7: # Weekly reset on Monday at midnight
|
||||
if value == 1: # Daily reset at the configured time of day
|
||||
return _next_occurrence(base_midnight, reset_time_of_day, current_time, timedelta(days=1))
|
||||
elif value == 7: # Weekly reset on Monday at the configured time of day
|
||||
days_until_monday = (7 - current_time.weekday()) % 7
|
||||
if days_until_monday == 0: # If today is Monday
|
||||
days_until_monday = 7
|
||||
return base_midnight + timedelta(days=days_until_monday)
|
||||
elif value == 30: # Monthly reset on 1st at midnight
|
||||
# Get 1st of next month at midnight
|
||||
if current_time.month == 12:
|
||||
next_reset = datetime(
|
||||
year=current_time.year + 1,
|
||||
month=1,
|
||||
day=1,
|
||||
hour=0,
|
||||
minute=0,
|
||||
second=0,
|
||||
microsecond=0,
|
||||
tzinfo=tz,
|
||||
)
|
||||
else:
|
||||
next_reset = datetime(
|
||||
year=current_time.year,
|
||||
month=current_time.month + 1,
|
||||
day=1,
|
||||
hour=0,
|
||||
minute=0,
|
||||
second=0,
|
||||
microsecond=0,
|
||||
tzinfo=tz,
|
||||
)
|
||||
return next_reset
|
||||
else: # Custom day value - next interval is value days from current
|
||||
return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=value)
|
||||
upcoming_monday = base_midnight + timedelta(days=days_until_monday)
|
||||
return _next_occurrence(upcoming_monday, reset_time_of_day, current_time, timedelta(days=7))
|
||||
elif value == 30: # Monthly reset on 1st at the configured time of day
|
||||
return _handle_month_reset(current_time, base_midnight, 1, reset_time_of_day)
|
||||
else: # Custom day value - next interval is value days from the start of today
|
||||
return _apply_time_of_day(base_midnight + timedelta(days=value), reset_time_of_day)
|
||||
|
||||
|
||||
def _handle_hour_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime:
|
||||
|
|
@ -316,36 +353,30 @@ def _handle_second_reset(current_time: datetime, base_midnight: datetime, value:
|
|||
return current_time.replace(hour=next_hour, minute=next_minute, second=next_second, microsecond=0)
|
||||
|
||||
|
||||
def _handle_month_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime:
|
||||
def _handle_month_reset(
|
||||
current_time: datetime,
|
||||
base_midnight: datetime,
|
||||
value: int,
|
||||
reset_time_of_day: time,
|
||||
) -> datetime:
|
||||
"""
|
||||
Handle monthly reset times. For monthly resets, we always reset at the start of the next month.
|
||||
Handle monthly reset times. Resets land on the 1st at `reset_time_of_day`; if the
|
||||
1st of the current month at that time has already passed, roll to the 1st of next month.
|
||||
|
||||
Args:
|
||||
current_time: Current datetime
|
||||
base_midnight: Midnight of current day
|
||||
value: Number of months (currently only supports 1 month resets)
|
||||
reset_time_of_day: Wall-clock time the reset lands on
|
||||
|
||||
Returns:
|
||||
datetime: First day of next month at midnight
|
||||
datetime: First day of the next reset month at `reset_time_of_day`
|
||||
"""
|
||||
if value != 1:
|
||||
raise ValueError("Monthly resets currently only support 1 month intervals")
|
||||
|
||||
# Get the first day of next month
|
||||
if current_time.month == 12:
|
||||
next_month = 1
|
||||
next_year = current_time.year + 1
|
||||
else:
|
||||
next_month = current_time.month + 1
|
||||
next_year = current_time.year
|
||||
|
||||
return datetime(
|
||||
year=next_year,
|
||||
month=next_month,
|
||||
day=1,
|
||||
hour=0,
|
||||
minute=0,
|
||||
second=0,
|
||||
microsecond=0,
|
||||
tzinfo=current_time.tzinfo,
|
||||
)
|
||||
first_of_this_month = base_midnight.replace(day=1)
|
||||
candidate = _apply_time_of_day(first_of_this_month, reset_time_of_day)
|
||||
if candidate <= current_time:
|
||||
return _apply_time_of_day(_first_of_next_month(first_of_this_month), reset_time_of_day)
|
||||
return candidate
|
||||
|
|
|
|||
|
|
@ -148,6 +148,15 @@ def _parse_url_destination_allowlist_entry(
|
|||
return _normalize_host(parsed.hostname), scheme, port
|
||||
|
||||
|
||||
def provider_url_destination_candidates(value: str) -> Tuple[str, ...]:
|
||||
return tuple(
|
||||
candidate
|
||||
for part in value.split(",")
|
||||
for candidate in (part.strip(), part.strip().split("/", 1)[1] if "/" in part.strip() else "")
|
||||
if candidate
|
||||
)
|
||||
|
||||
|
||||
def is_url_destination_allowed_by_host(url: str, allowed_hosts: List[str]) -> bool:
|
||||
"""Return True when a credential-bearing provider URL is admin-allowlisted.
|
||||
|
||||
|
|
|
|||
|
|
@ -480,10 +480,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"""
|
||||
Filter out unsupported fields from JSON schema for Anthropic's output_format API.
|
||||
|
||||
Anthropic's output_format doesn't support certain JSON schema properties:
|
||||
- maxItems/minItems: Not supported for array types
|
||||
- minimum/maximum: Not supported for numeric types
|
||||
- minLength/maxLength: Not supported for string types
|
||||
Anthropic's output_format doesn't support certain JSON schema properties.
|
||||
These are constraints that cannot be enforced by the constrained-decoding
|
||||
grammar Anthropic compiles the schema into, so the API rejects them with a
|
||||
400 ``invalid_request_error`` (e.g. "output_format.schema: For 'array' type,
|
||||
property 'uniqueItems' is not supported"):
|
||||
- maxItems/minItems/uniqueItems/contains/minContains/maxContains/prefixItems: array constraints
|
||||
- minimum/maximum/exclusiveMinimum/exclusiveMaximum/multipleOf: numeric constraints
|
||||
- minLength/maxLength: string constraints
|
||||
- minProperties/maxProperties/patternProperties/propertyNames: object constraints
|
||||
- dependentRequired/dependentSchemas/unevaluatedProperties: object constraints
|
||||
- if/then/else/not: conditional and negation keywords
|
||||
|
||||
``oneOf`` is also rejected ("Schema type 'oneOf' is not supported") and is
|
||||
rewritten to ``anyOf``, matching the Anthropic SDK. Unknown keywords are
|
||||
ignored by the API, so anything not listed here passes through untouched.
|
||||
|
||||
This mirrors the transformation done by the Anthropic Python SDK.
|
||||
See: https://platform.claude.com/docs/en/build-with-claude/structured-outputs#how-sdk-transformation-works
|
||||
|
|
@ -504,33 +515,53 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if not isinstance(schema, dict):
|
||||
return schema
|
||||
|
||||
# All numeric/string/array constraints not supported by Anthropic
|
||||
unsupported_fields = {
|
||||
"maxItems",
|
||||
"minItems", # array constraints
|
||||
"minimum",
|
||||
"maximum", # numeric constraints
|
||||
"exclusiveMinimum",
|
||||
"exclusiveMaximum", # numeric constraints
|
||||
"minLength",
|
||||
"maxLength", # string constraints
|
||||
}
|
||||
|
||||
# Build description additions from removed constraints
|
||||
constraint_descriptions: list = []
|
||||
constraint_labels = {
|
||||
"minItems": "minimum number of items: {}",
|
||||
"maxItems": "maximum number of items: {}",
|
||||
"uniqueItems": "all array items must be unique",
|
||||
"contains": "array must contain an item matching: {}",
|
||||
"minContains": "minimum number of matching items: {}",
|
||||
"maxContains": "maximum number of matching items: {}",
|
||||
"prefixItems": "leading items must match, in order: {}",
|
||||
"minimum": "minimum value: {}",
|
||||
"maximum": "maximum value: {}",
|
||||
"exclusiveMinimum": "exclusive minimum value: {}",
|
||||
"exclusiveMaximum": "exclusive maximum value: {}",
|
||||
"multipleOf": "must be a multiple of {}",
|
||||
"minLength": "minimum length: {}",
|
||||
"maxLength": "maximum length: {}",
|
||||
"minProperties": "minimum number of properties: {}",
|
||||
"maxProperties": "maximum number of properties: {}",
|
||||
"patternProperties": "properties whose names match each pattern must satisfy: {}",
|
||||
"propertyNames": "property names must satisfy: {}",
|
||||
"dependentRequired": "dependent required properties: {}",
|
||||
"dependentSchemas": "dependent schemas: {}",
|
||||
"unevaluatedProperties": "unevaluated properties must satisfy: {}",
|
||||
"if": "conditional (if): {}",
|
||||
"then": "conditional (then): {}",
|
||||
"else": "conditional (else): {}",
|
||||
"not": "must not match: {}",
|
||||
}
|
||||
for field in unsupported_fields:
|
||||
if field in schema:
|
||||
constraint_descriptions.append(constraint_labels[field].format(schema[field]))
|
||||
unsupported_fields = set(constraint_labels)
|
||||
|
||||
# Build description additions from removed constraints. Iterating
|
||||
# constraint_labels (not the set) keeps the note order deterministic across
|
||||
# processes, so identical requests serialize identically regardless of
|
||||
# PYTHONHASHSEED and stay cache-friendly.
|
||||
constraint_descriptions: list = []
|
||||
for field, label in constraint_labels.items():
|
||||
if field not in schema:
|
||||
continue
|
||||
value = schema[field]
|
||||
# A falsy boolean constraint (e.g. ``uniqueItems: false``) imposes no
|
||||
# real requirement, so don't add a misleading advisory note for it.
|
||||
if isinstance(value, bool) and not value:
|
||||
continue
|
||||
# Sub-schema constraints (e.g. ``contains``) are serialized as JSON so
|
||||
# the advisory note preserves what the constraint actually required,
|
||||
# instead of just noting that it existed.
|
||||
note_value = json.dumps(value) if isinstance(value, (dict, list)) else value
|
||||
constraint_descriptions.append(label.format(note_value))
|
||||
|
||||
result: Dict[str, Any] = {}
|
||||
|
||||
|
|
@ -557,11 +588,17 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
elif key == "$defs" and isinstance(value, dict):
|
||||
result[key] = {k: AnthropicConfig.filter_anthropic_output_schema(v) for k, v in value.items()}
|
||||
elif key == "anyOf" and isinstance(value, list):
|
||||
result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value]
|
||||
result["anyOf"] = result.get("anyOf", []) + [
|
||||
AnthropicConfig.filter_anthropic_output_schema(item) for item in value
|
||||
]
|
||||
elif key == "allOf" and isinstance(value, list):
|
||||
result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value]
|
||||
elif key == "oneOf" and isinstance(value, list):
|
||||
result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value]
|
||||
# Anthropic rejects oneOf ("Schema type 'oneOf' is not supported");
|
||||
# the Anthropic SDK rewrites it to anyOf, so do the same.
|
||||
result["anyOf"] = result.get("anyOf", []) + [
|
||||
AnthropicConfig.filter_anthropic_output_schema(item) for item in value
|
||||
]
|
||||
else:
|
||||
result[key] = value
|
||||
|
||||
|
|
|
|||
|
|
@ -895,9 +895,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if _tool_choice_value is not None:
|
||||
optional_params["tool_choice"] = _tool_choice_value
|
||||
if param == "parallel_tool_calls":
|
||||
disable_parallel = not value
|
||||
optional_params["_parallel_tool_use_config"] = {
|
||||
"tool_choice": {"disable_parallel_tool_use": disable_parallel}
|
||||
"tool_choice": {"type": "auto", "disable_parallel_tool_use": not value}
|
||||
}
|
||||
if param == "thinking":
|
||||
if (
|
||||
|
|
@ -1208,6 +1207,22 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
def _merge_parallel_tool_use_config(additional_request_params: dict, parallel_tool_use_config: dict) -> dict:
|
||||
merged_entries = {
|
||||
key: (
|
||||
{
|
||||
**value,
|
||||
**additional_request_params[key],
|
||||
**{k: v for k, v in value.items() if k != "type"},
|
||||
}
|
||||
if isinstance(additional_request_params.get(key), dict) and isinstance(value, dict)
|
||||
else value
|
||||
)
|
||||
for key, value in parallel_tool_use_config.items()
|
||||
}
|
||||
return {**additional_request_params, **merged_entries}
|
||||
|
||||
def _prepare_request_params(
|
||||
self, optional_params: dict, model: str, drop_params: bool = False
|
||||
) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]:
|
||||
|
|
@ -1276,15 +1291,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
# Handle parallel_tool_calls configuration
|
||||
parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None)
|
||||
if parallel_tool_use_config is not None and bedrock_converse_supports_parallel_tool_use_config(model):
|
||||
for key, value in parallel_tool_use_config.items():
|
||||
if (
|
||||
key in additional_request_params
|
||||
and isinstance(additional_request_params[key], dict)
|
||||
and isinstance(value, dict)
|
||||
):
|
||||
additional_request_params[key].update(value)
|
||||
else:
|
||||
additional_request_params[key] = value
|
||||
additional_request_params = self._merge_parallel_tool_use_config(
|
||||
additional_request_params, parallel_tool_use_config
|
||||
)
|
||||
|
||||
additional_request_params.pop("parallel_tool_calls", None)
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ import contextlib
|
|||
import json
|
||||
from typing import Any, Optional
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
||||
|
|
@ -16,6 +18,8 @@ from ..base_aws_llm import BaseAWSLLM
|
|||
from ..common_utils import BedrockError
|
||||
from .transformation import BedrockRealtimeConfig
|
||||
|
||||
_CLIENT_MODALITIES_ADAPTER: TypeAdapter["list[str] | None"] = TypeAdapter(list[str] | None)
|
||||
|
||||
|
||||
class BedrockRealtime(BaseAWSLLM):
|
||||
"""Handler for Bedrock Nova Sonic realtime speech-to-speech API."""
|
||||
|
|
@ -124,6 +128,9 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
|
||||
|
||||
await websocket.send_text(json.dumps(transformation_config.session_created_event(model, logging_obj)))
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect")
|
||||
|
||||
# Track state for transformation
|
||||
session_state = {
|
||||
"current_output_item_id": None,
|
||||
|
|
@ -143,6 +150,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
transformation_config,
|
||||
model,
|
||||
session_state,
|
||||
logging_obj,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -179,6 +187,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
transformation_config: BedrockRealtimeConfig,
|
||||
model: str,
|
||||
session_state: dict,
|
||||
logging_obj: LiteLLMLogging | None = None,
|
||||
):
|
||||
"""Forward messages from client WebSocket to Bedrock stream."""
|
||||
from aws_sdk_bedrock_runtime.models import (
|
||||
|
|
@ -210,6 +219,23 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
for bedrock_message in transformed_messages:
|
||||
await send_to_bedrock(bedrock_message)
|
||||
|
||||
if logging_obj is not None:
|
||||
client_message_type: str | None = None
|
||||
requested_modalities: list[str] | None = None
|
||||
with contextlib.suppress(Exception):
|
||||
parsed_client_message = json.loads(message)
|
||||
client_message_type = parsed_client_message.get("type")
|
||||
if client_message_type == "session.update":
|
||||
requested_modalities = _CLIENT_MODALITIES_ADAPTER.validate_python(
|
||||
parsed_client_message.get("session", {}).get("modalities")
|
||||
)
|
||||
if client_message_type == "session.update":
|
||||
await client_ws.send_text(
|
||||
json.dumps(
|
||||
transformation_config.session_updated_event(model, logging_obj, requested_modalities)
|
||||
)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(f"Client to Bedrock forwarding ended: {e}", exc_info=True)
|
||||
for close_message in transformation_config.session_close_messages():
|
||||
|
|
|
|||
|
|
@ -623,35 +623,42 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
verbose_logger.warning(f"Unknown message type: {message_type}")
|
||||
return []
|
||||
|
||||
def transform_session_start_event(
|
||||
def _session_object(
|
||||
self,
|
||||
event: dict,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> OpenAIRealtimeStreamSessionEvents:
|
||||
"""
|
||||
Transform Bedrock sessionStart event to OpenAI session.created.
|
||||
|
||||
Args:
|
||||
event: Bedrock sessionStart event
|
||||
model: Model ID
|
||||
logging_obj: Logging object
|
||||
|
||||
Returns:
|
||||
OpenAI session.created event
|
||||
"""
|
||||
verbose_logger.debug("Handling sessionStart")
|
||||
|
||||
modalities: list[str] | None = None,
|
||||
) -> OpenAIRealtimeStreamSession:
|
||||
session = OpenAIRealtimeStreamSession(
|
||||
id=logging_obj.litellm_trace_id,
|
||||
modalities=["text", "audio"],
|
||||
modalities=modalities if modalities is not None else ["text", "audio"],
|
||||
)
|
||||
if model is not None and isinstance(model, str):
|
||||
session["model"] = model
|
||||
return session
|
||||
|
||||
def session_created_event(
|
||||
self,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> OpenAIRealtimeStreamSessionEvents:
|
||||
"""Build the OpenAI session.created event for this realtime session."""
|
||||
return OpenAIRealtimeStreamSessionEvents(
|
||||
type="session.created",
|
||||
session=session,
|
||||
session=self._session_object(model, logging_obj),
|
||||
event_id=str(uuid.uuid4()),
|
||||
)
|
||||
|
||||
def session_updated_event(
|
||||
self,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
modalities: list[str] | None = None,
|
||||
) -> OpenAIRealtimeStreamSessionEvents:
|
||||
"""Build the OpenAI session.updated ack reflecting the client's requested modalities."""
|
||||
return OpenAIRealtimeStreamSessionEvents(
|
||||
type="session.updated",
|
||||
session=self._session_object(model, logging_obj, modalities),
|
||||
event_id=str(uuid.uuid4()),
|
||||
)
|
||||
|
||||
|
|
@ -1169,8 +1176,6 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
|
||||
# Route to appropriate transformation method
|
||||
if "sessionStart" in event:
|
||||
session_created = self.transform_session_start_event(event, model, logging_obj)
|
||||
returned_messages.append(session_created)
|
||||
session_configuration_request = json.dumps({"configured": True})
|
||||
|
||||
elif "contentStart" in event:
|
||||
|
|
|
|||
|
|
@ -2107,8 +2107,7 @@ class BaseLLMHTTPHandler:
|
|||
rust_messages_response = await self._maybe_rust_anthropic_messages(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
rust_stream_eligible=bool(stream) and not self._has_agentic_completion_hook(logging_obj),
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -2266,8 +2265,7 @@ class BaseLLMHTTPHandler:
|
|||
*,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
stream: bool,
|
||||
rust_stream_eligible: bool,
|
||||
has_agentic_hook: bool,
|
||||
model: str,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
|
|
@ -2279,7 +2277,7 @@ class BaseLLMHTTPHandler:
|
|||
return None
|
||||
if litellm_params.get("rust") is not True and not BaseLLMHTTPHandler._rust_env_enabled():
|
||||
return None
|
||||
if stream and not rust_stream_eligible:
|
||||
if has_agentic_hook:
|
||||
return None
|
||||
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
|
|
|||
|
|
@ -322,7 +322,7 @@ class HuggingFaceEmbedding(BaseLLM):
|
|||
task = get_hf_task_embedding_for_model(model=model, task_type=task_type, api_base=HF_HUB_URL)
|
||||
# print_verbose(f"{model}, {task}")
|
||||
embed_url = ""
|
||||
if "https" in model:
|
||||
if model.startswith(("http://", "https://")):
|
||||
embed_url = model
|
||||
elif api_base:
|
||||
embed_url = api_base
|
||||
|
|
|
|||
|
|
@ -316,25 +316,6 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
|
|||
|
||||
return data
|
||||
|
||||
def get_api_base(self, api_base: Optional[str], model: str) -> str:
|
||||
"""
|
||||
Get the API base for the Huggingface API.
|
||||
|
||||
Do not add the chat/embedding/rerank extension here. Let the handler do this.
|
||||
"""
|
||||
if "https" in model:
|
||||
completion_url = model
|
||||
elif api_base is not None:
|
||||
completion_url = api_base
|
||||
elif "HF_API_BASE" in os.environ:
|
||||
completion_url = os.getenv("HF_API_BASE", "")
|
||||
elif "HUGGINGFACE_API_BASE" in os.environ:
|
||||
completion_url = os.getenv("HUGGINGFACE_API_BASE", "")
|
||||
else:
|
||||
completion_url = f"https://api-inference.huggingface.co/models/{model}"
|
||||
|
||||
return completion_url
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ def completion(
|
|||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if "https" in model:
|
||||
if model.startswith(("http://", "https://")):
|
||||
completion_url = model
|
||||
elif api_base:
|
||||
completion_url = api_base
|
||||
|
|
@ -96,7 +96,7 @@ def embedding(
|
|||
encoding=None,
|
||||
):
|
||||
# Create completion URL
|
||||
if "https" in model:
|
||||
if model.startswith(("http://", "https://")):
|
||||
embeddings_url = model
|
||||
elif api_base:
|
||||
embeddings_url = f"{api_base}/v1/embeddings"
|
||||
|
|
|
|||
|
|
@ -143,7 +143,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
|
|||
raise SagemakerError(status_code=response.status_code, message=response.text)
|
||||
|
||||
custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True)
|
||||
completion_stream = custom_stream_decoder.iter_bytes(response.iter_bytes(chunk_size=1024))
|
||||
completion_stream = custom_stream_decoder.iter_bytes(response.iter_bytes())
|
||||
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
|
|
@ -189,7 +189,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
|
|||
raise SagemakerError(status_code=response.status_code, message=response.text)
|
||||
|
||||
custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True)
|
||||
completion_stream = custom_stream_decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024))
|
||||
completion_stream = custom_stream_decoder.aiter_bytes(response.aiter_bytes())
|
||||
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
|
|
|
|||
|
|
@ -200,23 +200,12 @@ class SagemakerLLM(BaseAWSLLM):
|
|||
# Add model_id as InferenceComponentName header
|
||||
# boto3 doc: https://docs.aws.amazon.com/sagemaker/latest/APIReference/API_runtime_InvokeEndpoint.html
|
||||
prepared_request.headers.update({"X-Amzn-SageMaker-Inference-Component": model_id})
|
||||
sync_handler = _get_httpx_client()
|
||||
sync_response = sync_handler.post(
|
||||
url=prepared_request.url,
|
||||
completion_stream = self.make_sync_call(
|
||||
api_base=prepared_request.url,
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
data=prepared_request.body,
|
||||
stream=stream,
|
||||
data=cast(str, prepared_request.body), # cast-ok: signed body is a JSON str, mirrors async path
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if sync_response.status_code != 200:
|
||||
raise SagemakerError(
|
||||
status_code=sync_response.status_code,
|
||||
message=str(sync_response.read()),
|
||||
)
|
||||
|
||||
decoder = AWSEventStreamDecoder(model="")
|
||||
|
||||
completion_stream = decoder.iter_bytes(sync_response.iter_bytes(chunk_size=1024))
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
|
|
@ -334,6 +323,29 @@ class SagemakerLLM(BaseAWSLLM):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def make_sync_call(
|
||||
self,
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: str,
|
||||
logging_obj,
|
||||
client=None,
|
||||
):
|
||||
if client is None:
|
||||
client = _get_httpx_client()
|
||||
sync_response = client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=data,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
if sync_response.status_code != 200:
|
||||
raise SagemakerError(status_code=sync_response.status_code, message=str(sync_response.read()))
|
||||
|
||||
decoder = AWSEventStreamDecoder(model="")
|
||||
return decoder.iter_bytes(sync_response.iter_bytes())
|
||||
|
||||
async def make_async_call(
|
||||
self,
|
||||
api_base: str,
|
||||
|
|
@ -358,7 +370,7 @@ class SagemakerLLM(BaseAWSLLM):
|
|||
raise SagemakerError(status_code=response.status_code, message=response.text)
|
||||
|
||||
decoder = AWSEventStreamDecoder(model="")
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024))
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes())
|
||||
|
||||
return completion_stream
|
||||
|
||||
|
|
|
|||
|
|
@ -5111,7 +5111,10 @@ def completion( # type: ignore
|
|||
try:
|
||||
if base_url is not None:
|
||||
api_base = base_url
|
||||
if num_retries is not None:
|
||||
is_router_call = any("model_group" in (kwargs.get(k) or ()) for k in ("metadata", "litellm_metadata"))
|
||||
if is_router_call:
|
||||
max_retries = 0
|
||||
elif num_retries is not None:
|
||||
max_retries = num_retries
|
||||
logging: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, litellm_logging_obj)
|
||||
fallbacks = fallbacks or litellm.model_fallbacks
|
||||
|
|
|
|||
|
|
@ -17566,6 +17566,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
|
|
@ -18233,6 +18288,60 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
|
|
@ -19585,6 +19694,63 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"rpm": 15,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 250000,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3-flash-preview": {
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -19691,6 +19857,63 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/pricing/gemini-3",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-omni-flash-preview": {
|
||||
"input_cost_per_audio_token": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
|
|
@ -19971,6 +20194,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"source": "https://ai.google.dev/pricing/gemini-3",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-2.5-pro-preview-tts": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
|
|
@ -37232,6 +37510,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import html as _html
|
|||
import json
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
|
@ -13,6 +14,7 @@ from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Resp
|
|||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -20,7 +22,9 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
normalize_token_endpoint_auth_method,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod
|
||||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
||||
_bridge_mint_error_response,
|
||||
_BridgeMintReady,
|
||||
|
|
@ -111,6 +115,9 @@ def encode_state_with_base_url(
|
|||
client_redirect_uri: Optional[str] = None,
|
||||
litellm_user_id: str | None = None,
|
||||
mcp_server_id: str | None = None,
|
||||
dcr_client_id: str | None = None,
|
||||
dcr_client_secret: str | None = None,
|
||||
dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Encode the base_url, original state, and PKCE parameters using encryption.
|
||||
|
|
@ -124,8 +131,18 @@ def encode_state_with_base_url(
|
|||
litellm_user_id: The SSO-authenticated litellm user captured at the bridge authorize
|
||||
(interactive dcr_bridge oauth_delegate only); the callback seals it into the gateway
|
||||
authorization code so the token mint can bind the envelope to this user
|
||||
mcp_server_id: The bridge server the interactive flow targets, sealed alongside
|
||||
litellm_user_id so the gateway code cannot be replayed against another server
|
||||
mcp_server_id: The server the flow targets, sealed alongside litellm_user_id (bridge) or
|
||||
dcr_client_id (ephemeral mint) so the gateway code cannot be replayed against another
|
||||
server
|
||||
dcr_client_id: The ephemeral DCR client the gateway minted at authorize for a
|
||||
client-forwarded-token server with no caller-supplied client; the callback seals it
|
||||
into the forwarded authorization code so the token exchange can authenticate with it
|
||||
while the gateway stores nothing
|
||||
dcr_client_secret: The minted client's secret, when the upstream issued one
|
||||
dcr_token_endpoint_auth_method: The token-endpoint auth method the upstream's registration
|
||||
response granted the minted client, sealed alongside the credentials so the exchange
|
||||
authenticates the way the upstream expects instead of falling back to the server row's
|
||||
configured method
|
||||
|
||||
Returns:
|
||||
An encrypted string that encodes all values
|
||||
|
|
@ -138,6 +155,9 @@ def encode_state_with_base_url(
|
|||
"client_redirect_uri": client_redirect_uri,
|
||||
"litellm_user_id": litellm_user_id,
|
||||
"mcp_server_id": mcp_server_id,
|
||||
"dcr_client_id": dcr_client_id,
|
||||
"dcr_client_secret": dcr_client_secret,
|
||||
"dcr_token_endpoint_auth_method": dcr_token_endpoint_auth_method,
|
||||
}
|
||||
state_json = json.dumps(state_data, sort_keys=True)
|
||||
encrypted_state = encrypt_value_helper(state_json)
|
||||
|
|
@ -217,6 +237,93 @@ def open_bridge_authorization_code(code: str) -> _BridgeAuthorizationCode | None
|
|||
return None
|
||||
|
||||
|
||||
_PASSTHROUGH_AUTH_CODE_PREFIX = "llm_ptcode_"
|
||||
|
||||
|
||||
class PassthroughAuthorizationCode(BaseModel):
|
||||
"""The ephemeral DCR client and upstream code the gateway seals into the authorization code it
|
||||
forwards for a client-forwarded-token server (``true_passthrough`` / ``oauth_delegate``) whose
|
||||
authorize fell through to gateway-side registration. These modes forbid the gateway from storing
|
||||
an OAuth client identity, so the minted client survives only inside this sealed value: the
|
||||
client echoes it back at the token endpoint, where the gateway recovers the client to
|
||||
authenticate the upstream exchange. ``mcp_server_id`` binds the code to the server it was minted
|
||||
for so it cannot be spent at another server's token endpoint."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
upstream_code: str = Field(min_length=1)
|
||||
client_id: str = Field(min_length=1)
|
||||
client_secret: str | None = None
|
||||
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None
|
||||
mcp_server_id: str = Field(min_length=1)
|
||||
|
||||
|
||||
def seal_passthrough_authorization_code(
|
||||
upstream_code: str,
|
||||
client_id: str,
|
||||
client_secret: str | None,
|
||||
mcp_server_id: str,
|
||||
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
|
||||
) -> str:
|
||||
"""Seal the upstream authorization code together with the ephemeral DCR client that authorized
|
||||
it. Encrypted with the same authenticated symmetric helper as the OAuth state and bridge codes,
|
||||
so the client can neither read the (possibly confidential) client credentials nor forge a
|
||||
code."""
|
||||
payload = json.dumps(
|
||||
{
|
||||
"upstream_code": upstream_code,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"token_endpoint_auth_method": token_endpoint_auth_method,
|
||||
"mcp_server_id": mcp_server_id,
|
||||
},
|
||||
sort_keys=True,
|
||||
)
|
||||
return _PASSTHROUGH_AUTH_CODE_PREFIX + encrypt_value_helper(payload)
|
||||
|
||||
|
||||
def open_passthrough_authorization_code(code: str) -> PassthroughAuthorizationCode | None:
|
||||
"""Recover the sealed ephemeral client and upstream code, or ``None`` when ``code`` is not a
|
||||
gateway passthrough code or does not decrypt / validate, so a raw upstream code falls through to
|
||||
the existing caller-supplied-client behavior."""
|
||||
if not code.startswith(_PASSTHROUGH_AUTH_CODE_PREFIX):
|
||||
return None
|
||||
decrypted = decrypt_value_helper(
|
||||
code[len(_PASSTHROUGH_AUTH_CODE_PREFIX) :], "passthrough_authorization_code", return_original_value=False
|
||||
)
|
||||
if not isinstance(decrypted, str):
|
||||
return None
|
||||
try:
|
||||
return PassthroughAuthorizationCode.model_validate_json(decrypted)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def redeem_passthrough_authorization_code(
|
||||
code: str | None, mcp_server: MCPServer, code_verifier: str | None
|
||||
) -> PassthroughAuthorizationCode | None:
|
||||
"""The single redemption gate for sealed passthrough codes: a raw or foreign code returns
|
||||
``None`` so the caller keeps its existing behavior, while a genuine sealed code must be spent
|
||||
at the server it was minted for and must carry the PKCE verifier of the S256 flow that minted
|
||||
it (the mint refuses downgraded flows, so a verifier-less redemption is an interception
|
||||
attempt, not a legitimate client)."""
|
||||
if not code:
|
||||
return None
|
||||
sealed = open_passthrough_authorization_code(code)
|
||||
if sealed is None:
|
||||
return None
|
||||
if sealed.mcp_server_id != mcp_server.server_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Authorization code was issued for a different MCP server",
|
||||
)
|
||||
if not code_verifier:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="code_verifier is required to redeem this authorization code",
|
||||
)
|
||||
return sealed
|
||||
|
||||
|
||||
def _redirect_to_litellm_login(request: Request) -> RedirectResponse:
|
||||
"""Send an unauthenticated browser through litellm login before the interactive bridge authorize
|
||||
can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code,
|
||||
|
|
@ -594,10 +701,18 @@ async def authorize_with_server(
|
|||
code_challenge_method: Optional[str] = None,
|
||||
response_type: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
ephemeral_dcr_client: "EphemeralDcrClient | None" = None,
|
||||
):
|
||||
_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")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
|
||||
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
|
||||
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
|
||||
),
|
||||
)
|
||||
|
||||
if mcp_server.is_dcr_bridge:
|
||||
# Enforce S256 PKCE on both bridge arms. The relay arm forwards the validated,
|
||||
|
|
@ -605,7 +720,10 @@ async def authorize_with_server(
|
|||
# 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):
|
||||
# A gateway-minted ephemeral client is registered against {base}/callback, so its
|
||||
# flow must run the short-circuit arm; the relay arm is only for clients that
|
||||
# registered themselves through the front door and hold their own redirect binding.
|
||||
if _dcr_bridge_relays_client_registration(mcp_server) and ephemeral_dcr_client is None:
|
||||
return _redirect_to_upstream_authorize(
|
||||
mcp_server=mcp_server,
|
||||
client_id=client_id,
|
||||
|
|
@ -649,7 +767,12 @@ async def authorize_with_server(
|
|||
code_challenge_method=code_challenge_method,
|
||||
client_redirect_uri=redirect_uri,
|
||||
litellm_user_id=litellm_user_id,
|
||||
mcp_server_id=mcp_server.server_id if litellm_user_id else None,
|
||||
mcp_server_id=mcp_server.server_id if (litellm_user_id or ephemeral_dcr_client) else None,
|
||||
dcr_client_id=ephemeral_dcr_client.client_id if ephemeral_dcr_client else None,
|
||||
dcr_client_secret=ephemeral_dcr_client.client_secret if ephemeral_dcr_client else None,
|
||||
dcr_token_endpoint_auth_method=ephemeral_dcr_client.token_endpoint_auth_method
|
||||
if ephemeral_dcr_client
|
||||
else None,
|
||||
)
|
||||
relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES)
|
||||
|
||||
|
|
@ -696,23 +819,40 @@ async def exchange_token_with_server(
|
|||
code_verifier: Optional[str],
|
||||
refresh_token: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
if grant_type not in ("authorization_code", "refresh_token"):
|
||||
raise HTTPException(status_code=400, detail="Unsupported grant_type")
|
||||
|
||||
if mcp_server.token_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server token url is not set")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"MCP server token url is not configured. Servers with no url (OpenAPI spec or "
|
||||
"stdio) run no resource discovery, so set Token URL manually, or set Issuer to "
|
||||
"discover it from the identity provider (RFC 8414)."
|
||||
),
|
||||
)
|
||||
|
||||
# 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.
|
||||
# The id, secret, and token-endpoint auth method 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. Symmetrically, a caller-side client (an ephemeral mint
|
||||
# recovered from a sealed code) must authenticate the way its own registration was granted,
|
||||
# not the way the server row is configured; callers that carry no method keep the row's method
|
||||
# as before.
|
||||
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_id else client_secret
|
||||
resolved_auth_method = (
|
||||
mcp_server.token_endpoint_auth_method
|
||||
if mcp_server.client_id
|
||||
else (client_token_endpoint_auth_method or mcp_server.token_endpoint_auth_method)
|
||||
)
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=mcp_server.token_endpoint_auth_method,
|
||||
auth_method=resolved_auth_method,
|
||||
client_id=resolved_client_id,
|
||||
client_secret=resolved_client_secret,
|
||||
)
|
||||
|
|
@ -1215,7 +1355,7 @@ async def _persist_dcr_client_registration(
|
|||
return "failed"
|
||||
|
||||
|
||||
def _client_supplied_redirect_uris(value: object) -> list[str] | None:
|
||||
def client_supplied_redirect_uris(value: object) -> list[str] | None:
|
||||
"""RFC 7591 redirect_uris must be a non-empty array of URI strings. Any other shape (not a list,
|
||||
an empty list, or a list holding a non-string or empty-string element) yields None so every
|
||||
register arm falls back to the gateway callback instead of echoing a malformed value back to the
|
||||
|
|
@ -1227,6 +1367,142 @@ def _client_supplied_redirect_uris(value: object) -> list[str] | None:
|
|||
return uris if len(uris) == len(value) else None
|
||||
|
||||
|
||||
async def _post_dcr_registration(
|
||||
registration_url: str,
|
||||
register_data: Mapping[str, object],
|
||||
server_id: str,
|
||||
) -> httpx.Response:
|
||||
"""POST an RFC 7591 registration to the upstream and return its response, relaying a classified
|
||||
upstream rejection instead of a generic 500 and failing loud on an absent response."""
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
registration_url,
|
||||
headers=headers,
|
||||
json=register_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
status_code, detail = dcr_fault_detail(classify_upstream_dcr_rejection(exc.response, log_context=server_id))
|
||||
raise HTTPException(status_code=status_code, detail=detail) from exc
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no response",
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
class EphemeralDcrClient(BaseModel):
|
||||
"""A DCR client minted for a single authorize round trip and never stored by the gateway."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
client_id: str = Field(min_length=1)
|
||||
client_secret: str | None = None
|
||||
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None
|
||||
|
||||
|
||||
_EPHEMERAL_DCR_CLIENT_CACHE = InMemoryCache(default_ttl=_OAUTH_STATE_COOKIE_TTL_SECONDS)
|
||||
_EPHEMERAL_DCR_MINT_LOCKS: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
async def mint_ephemeral_dcr_client(request: Request, mcp_server: MCPServer) -> EphemeralDcrClient | None:
|
||||
"""Mint a throwaway OAuth client via the upstream's RFC 7591 registration endpoint for a
|
||||
client-forwarded-token server whose authorize arrived with no client_id. Returns ``None`` when
|
||||
the upstream exposes no registration endpoint, so the caller keeps its existing failure path.
|
||||
The minted client is deliberately not persisted anywhere: ``true_passthrough`` /
|
||||
``oauth_delegate`` require the gateway to hold no OAuth client identity, so it survives only in
|
||||
the encrypted OAuth state and the sealed authorization code the callback forwards.
|
||||
|
||||
Reloading the authorize page or retrying a flow must not register a fresh upstream client every
|
||||
time (an OAuth client identifies the application, not the user, so reuse is semantically
|
||||
correct). A per-process TTL cache bounded to the OAuth state cookie's lifetime dedupes the mint
|
||||
per (server, gateway origin), and a per-server lock single-flights concurrent mints (the
|
||||
``_OAUTH_METADATA_FETCH_LOCKS`` pattern; keyed by server_id alone so the lock registry stays
|
||||
bounded by the server count even when the request origin varies) so parallel authorize requests
|
||||
cannot each register an upstream client; the cache stamps nothing onto the server record and
|
||||
correctness never depends on it because the sealed state carries the client through the flow."""
|
||||
if mcp_server.registration_url is None:
|
||||
return None
|
||||
request_base_url = get_request_base_url(request)
|
||||
cache_key = f"mcp_ephemeral_dcr_client:{mcp_server.server_id}:{request_base_url}"
|
||||
cached = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key)
|
||||
if isinstance(cached, EphemeralDcrClient):
|
||||
return cached
|
||||
lock = _EPHEMERAL_DCR_MINT_LOCKS.setdefault(mcp_server.server_id, asyncio.Lock())
|
||||
async with lock:
|
||||
cached_after_wait = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key)
|
||||
if isinstance(cached_after_wait, EphemeralDcrClient):
|
||||
return cached_after_wait
|
||||
register_data: dict[str, object] = {
|
||||
"client_name": mcp_server.server_name or mcp_server.server_id,
|
||||
"redirect_uris": [f"{request_base_url}/callback"],
|
||||
"grant_types": ["authorization_code", "refresh_token"],
|
||||
"response_types": ["code"],
|
||||
"token_endpoint_auth_method": "none",
|
||||
}
|
||||
response = await _post_dcr_registration(
|
||||
registration_url=mcp_server.registration_url,
|
||||
register_data=register_data,
|
||||
server_id=mcp_server.server_id,
|
||||
)
|
||||
try:
|
||||
registration = _DcrClientRegistration.model_validate_json(response.text)
|
||||
except ValidationError as exc:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no usable client_id",
|
||||
) from exc
|
||||
if not registration.client_id:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no usable client_id",
|
||||
)
|
||||
minted = EphemeralDcrClient(
|
||||
client_id=registration.client_id,
|
||||
client_secret=registration.client_secret,
|
||||
token_endpoint_auth_method=normalize_token_endpoint_auth_method(registration.token_endpoint_auth_method),
|
||||
)
|
||||
_EPHEMERAL_DCR_CLIENT_CACHE.set_cache(cache_key, minted)
|
||||
return minted
|
||||
|
||||
|
||||
async def resolve_ephemeral_dcr_client(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
code_challenge: str | None,
|
||||
code_challenge_method: str | None,
|
||||
redirect_uri: str,
|
||||
) -> EphemeralDcrClient | None:
|
||||
"""The single owner of the gateway-side mint policy for a clientless authorize. Returns
|
||||
``None`` for servers whose mode does not permit gateway minting and for upstreams without a
|
||||
registration endpoint, so those callers keep their existing failure paths: plain ``oauth2``
|
||||
keeps its persisted-client contract, and the interactive ``oauth_delegate`` dcr_bridge
|
||||
sign-in has its own sealed-identity flow. ``true_passthrough`` mints regardless of the
|
||||
``dcr_bridge`` flag (the UI creates passthrough servers with the flag on by default): a
|
||||
minted flow runs the bridge short-circuit arm, while the relay front door remains for
|
||||
external clients that registered themselves. Flows that could never succeed fail loud
|
||||
before any upstream registration: a missing ``authorization_url``, a downgraded PKCE pair
|
||||
(without S256 the sealed code would be bearer-redeemable by any authenticated caller who
|
||||
intercepts the redirect), or an untrusted ``redirect_uri`` (a rejected redirect must not be
|
||||
usable to generate orphan IdP clients)."""
|
||||
if not (mcp_server.is_true_passthrough or (mcp_server.is_oauth_delegate and not mcp_server.is_dcr_bridge)):
|
||||
return None
|
||||
if mcp_server.authorization_url is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="MCP server authorization url is not set",
|
||||
)
|
||||
_require_s256_pkce(code_challenge, code_challenge_method)
|
||||
validate_trusted_redirect_uri(request, redirect_uri)
|
||||
return await mint_ephemeral_dcr_client(request, mcp_server)
|
||||
|
||||
|
||||
async def register_client_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -1262,7 +1538,14 @@ async def register_client_with_server(
|
|||
return dummy_return
|
||||
|
||||
if mcp_server.authorization_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server authorization url is not set")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
|
||||
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
|
||||
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
|
||||
),
|
||||
)
|
||||
|
||||
if mcp_server.registration_url is None:
|
||||
return dummy_return
|
||||
|
|
@ -1281,30 +1564,11 @@ async def register_client_with_server(
|
|||
"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",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
mcp_server.registration_url,
|
||||
headers=headers,
|
||||
json=register_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
status_code, detail = dcr_fault_detail(
|
||||
classify_upstream_dcr_rejection(exc.response, log_context=mcp_server.server_id)
|
||||
)
|
||||
raise HTTPException(status_code=status_code, detail=detail) from exc
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no response",
|
||||
)
|
||||
response = await _post_dcr_registration(
|
||||
registration_url=mcp_server.registration_url,
|
||||
register_data=register_data,
|
||||
server_id=mcp_server.server_id,
|
||||
)
|
||||
|
||||
token_response = response.json()
|
||||
|
||||
|
|
@ -1542,11 +1806,23 @@ async def callback(
|
|||
# envelope to this user. Every other flow forwards the raw code unchanged.
|
||||
litellm_user_id = state_data.get("litellm_user_id")
|
||||
mcp_server_id = state_data.get("mcp_server_id")
|
||||
dcr_client_id = state_data.get("dcr_client_id")
|
||||
dcr_client_secret = state_data.get("dcr_client_secret")
|
||||
forwarded_code = code
|
||||
if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id:
|
||||
forwarded_code = seal_bridge_authorization_code(
|
||||
upstream_code=code, litellm_user_id=litellm_user_id, mcp_server_id=mcp_server_id
|
||||
)
|
||||
elif isinstance(dcr_client_id, str) and dcr_client_id and isinstance(mcp_server_id, str) and mcp_server_id:
|
||||
forwarded_code = seal_passthrough_authorization_code(
|
||||
upstream_code=code,
|
||||
client_id=dcr_client_id,
|
||||
client_secret=dcr_client_secret if isinstance(dcr_client_secret, str) and dcr_client_secret else None,
|
||||
mcp_server_id=mcp_server_id,
|
||||
token_endpoint_auth_method=normalize_token_endpoint_auth_method(
|
||||
state_data.get("dcr_token_endpoint_auth_method")
|
||||
),
|
||||
)
|
||||
|
||||
params = {"code": forwarded_code, "state": original_state}
|
||||
complete_returned_url = _append_query_params(redirect_uri, params)
|
||||
|
|
@ -2137,7 +2413,7 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
|
|||
|
||||
request_data = await _read_request_body(request=request)
|
||||
data: dict = {**request_data}
|
||||
client_redirect_uris = _client_supplied_redirect_uris(data.get("redirect_uris"))
|
||||
client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris"))
|
||||
|
||||
dummy_return = {
|
||||
"client_id": mcp_server_name or "dummy_client",
|
||||
|
|
|
|||
|
|
@ -224,6 +224,20 @@ def _uses_issuer_anchor(manual_issuer: str | None, is_discovery_auth_type: bool)
|
|||
return _blank_to_none(manual_issuer) is not None and is_discovery_auth_type
|
||||
|
||||
|
||||
def _has_oauth_discovery_source(server_url: str | None, use_issuer_anchor: bool) -> bool:
|
||||
"""Whether the server has any source OAuth discovery can fetch metadata from.
|
||||
|
||||
Resource-rooted discovery (RFC 9728) is fetched from the server ``url``, so spec-only
|
||||
(OpenAPI) and stdio servers, which have none, could never discover: their OAuth endpoints
|
||||
stayed unset unless entered manually and ``/authorize`` served its 400 with no hint of why.
|
||||
An admin-pinned issuer is a trust anchor in its own right (RFC 8414 section 3.3) whose
|
||||
metadata fetch does not touch the resource at all, so an anchored server can discover with
|
||||
no ``url``. Called by both build paths (config and DB) so the two cannot disagree on when
|
||||
discovery is reachable.
|
||||
"""
|
||||
return bool(server_url) or use_issuer_anchor
|
||||
|
||||
|
||||
def _endpoints_yield_to_issuer(
|
||||
issuer: str | None,
|
||||
is_discovery_auth_type: bool,
|
||||
|
|
@ -610,6 +624,34 @@ def _passthrough_token_from_mcp_auth_header(
|
|||
return None
|
||||
|
||||
|
||||
async def _materialize_auth_headers(auth: httpx.Auth | None) -> dict[str, str] | None:
|
||||
"""Extract the header a resolved ``httpx.Auth`` would set, as a plain dict, or None.
|
||||
|
||||
OpenAPI tool closures egress through ``AsyncHTTPHandler`` methods that accept headers but no
|
||||
``auth``, so a resolved credential must be materialized into a header value. Driving one step
|
||||
of the auth's own flow (against a throwaway request that is never sent) keeps this generic
|
||||
across every auth shape without per-class branching; ``header_name`` is the resolver-arm
|
||||
convention for "this auth sets a header" (``NoOpAuth`` has none and yields nothing to apply).
|
||||
The materialized value is point-in-time: flow behaviors past the first request, like the M2M
|
||||
one-shot 401 refetch, do not apply on this arm.
|
||||
"""
|
||||
if auth is None:
|
||||
return None
|
||||
header_name = getattr(auth, "header_name", None)
|
||||
if not isinstance(header_name, str) or not header_name:
|
||||
return None
|
||||
probe = httpx.Request("GET", "http://localhost/")
|
||||
flow = auth.async_auth_flow(probe)
|
||||
try:
|
||||
first_request = await flow.__anext__()
|
||||
except StopAsyncIteration:
|
||||
return None
|
||||
finally:
|
||||
await flow.aclose()
|
||||
header_value = first_request.headers.get(header_name)
|
||||
return {header_name: header_value} if header_value else 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
|
||||
|
|
@ -1226,7 +1268,12 @@ class MCPServerManager:
|
|||
manual_token_url = _blank_to_none(server_config.get("token_url"))
|
||||
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
|
||||
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
|
||||
obo_needs_discovery = self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
server_config.get("token_exchange_endpoint"),
|
||||
manual_token_url,
|
||||
)
|
||||
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type or obo_needs_discovery)
|
||||
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
|
||||
manual_issuer,
|
||||
is_discovery_auth_type,
|
||||
|
|
@ -1234,17 +1281,12 @@ class MCPServerManager:
|
|||
manual_token_url,
|
||||
manual_registration_url,
|
||||
)
|
||||
should_discover = bool(server_url) and (
|
||||
is_discovery_auth_type
|
||||
or self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
server_config.get("token_exchange_endpoint"),
|
||||
manual_token_url,
|
||||
)
|
||||
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
|
||||
is_discovery_auth_type or obo_needs_discovery
|
||||
)
|
||||
if not should_discover:
|
||||
mcp_oauth_metadata = None
|
||||
elif manual_issuer is not None and is_discovery_auth_type:
|
||||
elif use_issuer_anchor and manual_issuer is not None:
|
||||
mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url)
|
||||
else:
|
||||
mcp_oauth_metadata = await self._descovery_metadata(
|
||||
|
|
@ -1640,7 +1682,7 @@ class MCPServerManager:
|
|||
token_exchange_endpoint: Optional[str],
|
||||
) -> Optional[MCPOAuthMetadata]:
|
||||
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
|
||||
needs_discovery = bool(server_url) and (
|
||||
needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
|
||||
(is_discovery_auth_type and not has_all_upstream_oauth_fields)
|
||||
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
|
||||
)
|
||||
|
|
@ -1759,13 +1801,17 @@ class MCPServerManager:
|
|||
manual_token_url = _blank_to_none(mcp_server.token_url)
|
||||
manual_registration_url = _blank_to_none(mcp_server.registration_url)
|
||||
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
|
||||
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
|
||||
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
|
||||
)
|
||||
token_exchange_endpoint = mcp_server.token_exchange_endpoint or (
|
||||
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None
|
||||
)
|
||||
use_issuer_anchor = _uses_issuer_anchor(
|
||||
manual_issuer,
|
||||
is_discovery_auth_type
|
||||
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
|
||||
)
|
||||
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
|
||||
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
|
||||
)
|
||||
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
|
||||
mcp_server=mcp_server,
|
||||
auth_type=auth_type,
|
||||
|
|
@ -1943,7 +1989,7 @@ class MCPServerManager:
|
|||
family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
|
||||
the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
|
||||
calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
|
||||
serves 400 "authorization url is not set" from /authorize until a later rebuild succeeds.
|
||||
serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
|
||||
Only fills row fields that are currently empty, never persists origin-fallback guesses
|
||||
(RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
|
||||
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
|
||||
|
|
@ -4705,6 +4751,61 @@ class MCPServerManager:
|
|||
)
|
||||
return oauth2_headers
|
||||
|
||||
async def resolve_openapi_upstream_auth(
|
||||
self,
|
||||
*,
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
raw_headers: dict[str, str] | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
forwarded_headers: dict[str, str] | None,
|
||||
) -> tuple[dict[str, str] | None, dict[str, str] | None]:
|
||||
"""Resolve the gateway-owned upstream credential for a spec_path (OpenAPI) tool call.
|
||||
|
||||
OpenAPI tools egress through a plain httpx call assembled from ContextVars, never through
|
||||
``_create_mcp_client``, so the v2 resolver graft there does not run for them and a resolved
|
||||
credential (authorization_code's stored per-user token, client_credentials' minted M2M
|
||||
token, token_exchange's exchanged token, passthrough's forwarded caller token) must be
|
||||
materialized into headers here. Returns ``(resolved_auth_headers, forwarded_headers)``:
|
||||
the resolved headers are authoritative over every other Authorization source (the same
|
||||
rule ``_resolve_v2_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
|
||||
back with any header the resolver claimed already dropped. Unmigrated (v1) servers resolve
|
||||
through the stored-token lookup instead, and a missing per-user credential raises the same
|
||||
discovery challenge the MCPClient path serves, rather than egressing unauthenticated.
|
||||
|
||||
The resolved headers carry only credentials the gateway itself resolved (a stored per-user
|
||||
token, a minted or exchanged token). Caller-supplied ``oauth2_headers`` are never promoted
|
||||
into them: on the v2 arm they feed only subject-token extraction (the designed RFC 8693
|
||||
input), and on the v1 arm their presence disables the stored lookup entirely, so a
|
||||
caller's gateway credential can never displace a per-server BYOK header or leak upstream
|
||||
as the resolved credential.
|
||||
"""
|
||||
spec = to_server_spec(mcp_server)
|
||||
if spec is None:
|
||||
if oauth2_headers:
|
||||
return None, forwarded_headers
|
||||
stored_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, None, user_api_key_auth)
|
||||
return stored_headers, forwarded_headers
|
||||
|
||||
subject_token: str | None = None
|
||||
if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
subject_token = self._extract_bearer_token(oauth2_headers, raw_headers)
|
||||
elif isinstance(spec.config, PassthroughConfig):
|
||||
inbound_token, forwarded_headers = _take_forwarded_authorization(forwarded_headers)
|
||||
per_server_token = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
subject_token = per_server_token if per_server_token is not None else inbound_token
|
||||
|
||||
resolved_auth, forwarded_headers = await self._resolve_v2_auth(
|
||||
server=mcp_server,
|
||||
spec=spec,
|
||||
provider=self._cred_provider,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=forwarded_headers,
|
||||
)
|
||||
return await _materialize_auth_headers(resolved_auth), forwarded_headers
|
||||
|
||||
async def _gather_openapi_tool_tasks(
|
||||
self,
|
||||
tasks: list[Any],
|
||||
|
|
@ -4796,6 +4897,7 @@ class MCPServerManager:
|
|||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
caller_oauth2_headers = oauth2_headers
|
||||
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth)
|
||||
|
||||
# For OpenAPI servers, call the tool handler directly instead of via MCP client
|
||||
|
|
@ -4813,22 +4915,32 @@ class MCPServerManager:
|
|||
auth_header_value = (
|
||||
_format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
|
||||
)
|
||||
forwarded_headers = _openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth)
|
||||
resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=caller_oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth),
|
||||
)
|
||||
|
||||
async def _call_openapi_via_handler():
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
)
|
||||
|
||||
auth_token = _request_auth_header.set(auth_header_value)
|
||||
extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -62,6 +62,14 @@ _request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = conte
|
|||
"_request_extra_headers", default=None
|
||||
)
|
||||
|
||||
# Per-request headers carrying the gateway-resolved upstream credential
|
||||
# (stored per-user OAuth token, minted M2M token, exchanged OBO token).
|
||||
# Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative
|
||||
# over every other Authorization source in _merge_openapi_tool_request_headers.
|
||||
_request_resolved_auth_headers: contextvars.ContextVar[dict[str, str] | None] = contextvars.ContextVar(
|
||||
"_request_resolved_auth_headers", default=None
|
||||
)
|
||||
|
||||
|
||||
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
|
||||
"""Ensure path params cannot introduce directory traversal."""
|
||||
|
|
@ -294,10 +302,15 @@ def _merge_openapi_tool_request_headers(
|
|||
"""Merge static closure headers with per-request ContextVar overrides.
|
||||
|
||||
Precedence (highest to lowest):
|
||||
1. ``_request_auth_header`` — BYOK override of ``Authorization``
|
||||
2. ``static_headers`` — operator-configured headers baked into the
|
||||
1. ``_request_resolved_auth_headers`` — the gateway-resolved upstream
|
||||
credential (stored per-user OAuth token, minted M2M token,
|
||||
exchanged OBO token). The resolver is authoritative: a BYOK or
|
||||
forwarded ``Authorization`` must not shadow it, mirroring
|
||||
``_resolve_v2_auth`` on the MCPClient path
|
||||
2. ``_request_auth_header`` — BYOK override of ``Authorization``
|
||||
3. ``static_headers`` — operator-configured headers baked into the
|
||||
tool closure at registration time
|
||||
3. ``_request_extra_headers`` — per-request headers forwarded from
|
||||
4. ``_request_extra_headers`` — per-request headers forwarded from
|
||||
the MCP caller (allowlisted by ``MCPServer.extra_headers``)
|
||||
|
||||
This matches the existing MCP invariant in
|
||||
|
|
@ -323,6 +336,12 @@ def _merge_openapi_tool_request_headers(
|
|||
del effective_headers[existing]
|
||||
effective_headers["Authorization"] = override_auth
|
||||
|
||||
resolved_auth_headers = _request_resolved_auth_headers.get() or {}
|
||||
for name, value in resolved_auth_headers.items():
|
||||
for existing in [k for k in effective_headers if k.lower() == name.lower()]:
|
||||
del effective_headers[existing]
|
||||
effective_headers[name] = value
|
||||
|
||||
return effective_headers
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,214 @@
|
|||
"""Store for the enterprise IdP identity assertion captured at SSO login (EMA).
|
||||
|
||||
The ``oauth2_id_jag`` egress arm needs the user's IdP ``id_token`` as its RFC 8693
|
||||
``subject_token``. A front-door client holds an identity-only ``llm_session_`` bearer, not an
|
||||
IdP assertion, so the assertion captured at the one SSO login is the only usable subject
|
||||
source for it. This module owns both sides of that state: the SSO callback persists here
|
||||
(write-through to the DB so a login on one pod is visible to every pod) and the resolver
|
||||
seam reads back by ``user_id``. Retention is gated on an ``oauth2_id_jag`` server actually
|
||||
being registered, so a gateway with no EMA upstream never stores bearer material.
|
||||
|
||||
The row is one encrypted payload per user, latest login wins. ``expires_at`` mirrors the
|
||||
id_token ``exp`` claim and is judged by the reader, never enforced by deletion here: an
|
||||
expired assertion with a refresh token is still renewable, and the DB row is the source of
|
||||
truth, the same contract as the per-user OAuth credential store.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import jwt
|
||||
from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_ASSERTION_DECRYPT_LOG_KEY = "sso_identity_assertion"
|
||||
_STR_ADAPTER: TypeAdapter[str] = TypeAdapter(str)
|
||||
_MAYBE_STR_ADAPTER: TypeAdapter[str | None] = TypeAdapter(str | None)
|
||||
|
||||
|
||||
class SSOIdentityAssertion(BaseModel):
|
||||
"""The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token,
|
||||
``expires_at`` bounds its usefulness, and the refresh token renews it without re-login."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
id_token: SecretStr
|
||||
refresh_token: SecretStr | None = None
|
||||
issuer: str | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
|
||||
class _IdTokenClaims(BaseModel):
|
||||
exp: float | None = None
|
||||
iss: str | None = None
|
||||
|
||||
|
||||
class _StoredAssertionPayload(BaseModel):
|
||||
id_token: str
|
||||
refresh_token: str | None = None
|
||||
issuer: str | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
|
||||
def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIdentityAssertion | None:
|
||||
"""The typed carrier built where the raw token response exists; ``None`` when the provider
|
||||
sent no id_token or sent one that is not a decodable JWT, since neither is exchangeable
|
||||
under EMA. Inputs are ``object`` because they come straight from the provider's untyped
|
||||
token response; this is the one boundary that validates them. The token arrived over TLS
|
||||
from the IdP's own token endpoint, so claims are read without signature verification,
|
||||
matching how the SSO callback already decodes it for identity."""
|
||||
raw_id_token = id_token if isinstance(id_token, str) and id_token else None
|
||||
if raw_id_token is None:
|
||||
return None
|
||||
raw_refresh_token = refresh_token if isinstance(refresh_token, str) and refresh_token else None
|
||||
try:
|
||||
claims = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False}))
|
||||
expires_at = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None
|
||||
except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login
|
||||
verbose_proxy_logger.warning(
|
||||
"SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress."
|
||||
)
|
||||
return None
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(raw_id_token),
|
||||
refresh_token=SecretStr(raw_refresh_token) if raw_refresh_token else None,
|
||||
issuer=claims.iss,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
|
||||
async def ema_assertion_retention_enabled() -> bool:
|
||||
"""Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
|
||||
retains bearer material while an EMA upstream exists to spend it on. Judged against the two
|
||||
configuration authorities: the pod-local config declaration and the shared DB row. The
|
||||
in-memory registry is deliberately not consulted in either direction; it is a per-process
|
||||
snapshot of the DB state that can be stale both ways (a server added on another pod would
|
||||
silently drop the write, one removed on another pod would keep retaining bearer material),
|
||||
and a gate guarding a shared-DB write must judge against that storage's authority."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global
|
||||
|
||||
config_servers = global_mcp_server_manager.config_mcp_servers.values()
|
||||
if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers):
|
||||
return True
|
||||
if prisma_client is None:
|
||||
return False
|
||||
row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
|
||||
return row is not None
|
||||
|
||||
|
||||
async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
payload: dict[str, str] = {
|
||||
"id_token": assertion.id_token.get_secret_value(),
|
||||
**({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}),
|
||||
**({"issuer": assertion.issuer} if assertion.issuer else {}),
|
||||
**({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}),
|
||||
}
|
||||
encoded = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload)))
|
||||
await prisma_client.db.litellm_ssoidentityassertion.upsert(
|
||||
where={"user_id": user_id},
|
||||
data={
|
||||
"create": {"user_id": user_id, "assertion_b64": encoded},
|
||||
"update": {"assertion_b64": encoded},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | None:
|
||||
"""The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key
|
||||
rotation), or unparseable. Expiry is not judged here; the reader owns that policy."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
row = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id})
|
||||
if row is None:
|
||||
return None
|
||||
raw = _MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
payload = _StoredAssertionPayload.model_validate_json(raw)
|
||||
except ValidationError:
|
||||
verbose_proxy_logger.warning(
|
||||
"Stored SSO identity assertion for user_id=%s could not be parsed; treating as absent.", user_id
|
||||
)
|
||||
return None
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(payload.id_token),
|
||||
refresh_token=SecretStr(payload.refresh_token) if payload.refresh_token else None,
|
||||
issuer=payload.issuer,
|
||||
expires_at=payload.expires_at,
|
||||
)
|
||||
|
||||
|
||||
async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
|
||||
"""Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
|
||||
mirroring the sibling per-user credential tables; an unreadable row is skipped so one
|
||||
corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop
|
||||
so the whole table's plaintext is never held in memory at once."""
|
||||
from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
async def _rotate_row(row: AssertionRow) -> bool:
|
||||
plaintext = _MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
)
|
||||
if plaintext is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping",
|
||||
row.user_id,
|
||||
)
|
||||
return False
|
||||
re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key))
|
||||
await prisma_client.db.litellm_ssoidentityassertion.update(
|
||||
where={"user_id": row.user_id},
|
||||
data={"assertion_b64": re_encrypted},
|
||||
)
|
||||
return True
|
||||
|
||||
rows = await prisma_client.db.litellm_ssoidentityassertion.find_many()
|
||||
outcomes = [await _rotate_row(row) for row in rows]
|
||||
verbose_proxy_logger.info(
|
||||
"rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",
|
||||
sum(outcomes),
|
||||
len(outcomes) - sum(outcomes),
|
||||
)
|
||||
|
||||
|
||||
async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIdentityAssertion | None) -> None:
|
||||
"""The SSO-callback hook: a no-op unless there is material AND an EMA server is registered.
|
||||
A store failure is logged and swallowed because the login itself must not fail on an
|
||||
egress-side write; the cost of a miss is a 401 challenge at the EMA upstream, not a lockout."""
|
||||
if assertion is None:
|
||||
return
|
||||
try:
|
||||
if not await ema_assertion_retention_enabled():
|
||||
return
|
||||
await persist_sso_identity_assertion(user_id, assertion)
|
||||
except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to persist the SSO identity assertion for EMA egress (user_id=%s): %s", user_id, exc
|
||||
)
|
||||
|
|
@ -376,6 +376,7 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
|
|
@ -2785,13 +2786,29 @@ if MCP_AVAILABLE:
|
|||
forwarded_headers = {}
|
||||
forwarded_headers[header_name] = value
|
||||
|
||||
resolved_auth_headers: dict[str, str] | None = None
|
||||
if mcp_server:
|
||||
(
|
||||
resolved_auth_headers,
|
||||
forwarded_headers,
|
||||
) = await global_mcp_server_manager.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=forwarded_headers,
|
||||
)
|
||||
|
||||
_auth_token = _request_auth_header.set(auth_header_value)
|
||||
_extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
_resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
local_content = await _handle_local_mcp_tool(name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(_auth_token)
|
||||
_request_extra_headers.reset(_extra_token)
|
||||
_request_resolved_auth_headers.reset(_resolved_token)
|
||||
response = CallToolResult(content=cast(Any, local_content), isError=False)
|
||||
|
||||
# Try managed MCP server tool (pass the full prefixed name)
|
||||
|
|
|
|||
|
|
@ -1856,6 +1856,17 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
default_team_member_models: Optional[List[str]] = None # default allowed_models seeded onto new team members
|
||||
|
||||
|
||||
class PatchTeamRequest(UpdateTeamRequest):
|
||||
"""
|
||||
Body of PATCH /team/{team_id}.
|
||||
|
||||
Identical to UpdateTeamRequest except team_id is optional, because PATCH takes it
|
||||
from the path. A team_id in the body is still accepted when it matches the path.
|
||||
"""
|
||||
|
||||
team_id: str | None = None
|
||||
|
||||
|
||||
class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
|
||||
"""
|
||||
internal type used to reset the budget on a team
|
||||
|
|
@ -2297,6 +2308,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="max response size in MB, if a response is larger than this size it will be rejected",
|
||||
)
|
||||
proxy_config_reload_interval_seconds: int = Field(
|
||||
30,
|
||||
gt=0,
|
||||
description="how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup",
|
||||
)
|
||||
cancel_on_disconnect: Optional[bool] = Field(
|
||||
None,
|
||||
description="cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure",
|
||||
|
|
@ -2762,6 +2778,30 @@ class LiteLLM_OrganizationTableUpdate(LiteLLM_BudgetTable):
|
|||
return values
|
||||
|
||||
|
||||
class OrganizationUpdateRequestV2(LiteLLMPydanticObjectBase):
|
||||
"""
|
||||
Typed PATCH body for ``/v2/organization/{organization_id}`` (RFC 7396 merge-patch).
|
||||
|
||||
Presence is read from ``model_fields_set``, so a sent field is written and an omitted one is
|
||||
left untouched. ``extra="forbid"`` makes an unknown key a 422 rather than a silent no-op, since
|
||||
the contract hinges on which keys are present. See the endpoint for the per-field clear tokens.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
organization_alias: str | None = None
|
||||
models: list[str] | None = None
|
||||
metadata: dict | None = None
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
max_budget: float | None = None
|
||||
soft_budget: float | None = None
|
||||
max_parallel_requests: int | None = None
|
||||
model_max_budget: dict | None = None
|
||||
budget_duration: str | None = None
|
||||
object_permission: LiteLLM_ObjectPermissionBase | None = None
|
||||
|
||||
|
||||
from litellm.models.organization import ( # noqa: E402
|
||||
LiteLLM_OrganizationTable as LiteLLM_OrganizationTable,
|
||||
)
|
||||
|
|
@ -4047,6 +4087,7 @@ class JWTAuthBuilderResult(TypedDict):
|
|||
token: str
|
||||
team_id: Optional[str]
|
||||
user_id: Optional[str]
|
||||
user_email: str | None
|
||||
end_user_id: Optional[str]
|
||||
org_id: Optional[str]
|
||||
team_membership: Optional[LiteLLM_TeamMembership]
|
||||
|
|
|
|||
|
|
@ -7,23 +7,46 @@ the base; specific fields are replaced so all traffic flows through the proxy
|
|||
and uses LiteLLM auth.
|
||||
"""
|
||||
|
||||
import re
|
||||
from copy import deepcopy
|
||||
from typing import Any, Dict, List, Mapping
|
||||
from typing import Any, Dict, List, Literal, Mapping
|
||||
|
||||
SupportedA2AVersion = Literal["0.3", "1.0"]
|
||||
|
||||
# Protocol versions LiteLLM can serve to A2A clients. The admin pins one per agent;
|
||||
# responses are normalized to it regardless of the upstream agent's own version.
|
||||
SUPPORTED_A2A_PROTOCOL_VERSIONS = ("0.3", "1.0")
|
||||
SUPPORTED_A2A_PROTOCOL_VERSIONS: tuple[SupportedA2AVersion, ...] = ("0.3", "1.0")
|
||||
|
||||
# Default served version when the agent card does not pin one.
|
||||
LITELLM_A2A_PROTOCOL_VERSION = "1.0"
|
||||
|
||||
|
||||
_PROTOCOL_VERSION_PATTERN = re.compile(
|
||||
r"^(\d+\.\d+)(?:\.\d+(?:-[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?)?$"
|
||||
)
|
||||
|
||||
|
||||
def normalize_protocol_version(version: object) -> SupportedA2AVersion | None:
|
||||
"""Map a raw ``protocolVersion`` value to the supported canonical major.minor version.
|
||||
|
||||
Accepts the bare major.minor convention of the 1.0 spec (``"0.3"``, ``"1.0"``) and the
|
||||
full semver forms older SDKs emit (``"0.3.0"``, ``"1.0.1"``, including prerelease and
|
||||
build suffixes like ``"0.3.0-rc1"``). Malformed strings, versions outside the
|
||||
supported set, and non-strings yield ``None``.
|
||||
"""
|
||||
if not isinstance(version, str):
|
||||
return None
|
||||
match = _PROTOCOL_VERSION_PATTERN.match(version)
|
||||
if match is None:
|
||||
return None
|
||||
major_minor = match.group(1)
|
||||
return next((supported for supported in SUPPORTED_A2A_PROTOCOL_VERSIONS if supported == major_minor), None)
|
||||
|
||||
|
||||
def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str:
|
||||
"""Return the validated protocol version an agent card pins, else the default."""
|
||||
version = card.get("protocolVersion") if card else None
|
||||
if version in SUPPORTED_A2A_PROTOCOL_VERSIONS:
|
||||
return version
|
||||
return LITELLM_A2A_PROTOCOL_VERSION
|
||||
normalized = normalize_protocol_version(card.get("protocolVersion") if card else None)
|
||||
return normalized if normalized is not None else LITELLM_A2A_PROTOCOL_VERSION
|
||||
|
||||
|
||||
# Security scheme exposed by the LiteLLM-fronted agent card. Always replaces
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from typing import Callable, Literal, Union
|
|||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.a2a.agent_card import normalize_protocol_version
|
||||
|
||||
A2AVersion = Literal["0.3", "1.0"]
|
||||
RequestId = Union[str, int, None]
|
||||
|
|
@ -103,16 +104,14 @@ def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: st
|
|||
def _detect_card_version(card: JsonDict) -> A2AVersion:
|
||||
"""Infer the wire version of an agent card dict.
|
||||
|
||||
``protocolVersion`` is the authoritative indicator; fall back to presence of
|
||||
``supportedInterfaces`` (a 1.0-only field) only when the explicit field is absent.
|
||||
Cards that set ``protocolVersion: "0.3"`` or carry neither signal are treated as 0.3.
|
||||
``protocolVersion`` is the authoritative indicator; semver values normalize to
|
||||
their major.minor (``"0.3.0"`` -> ``"0.3"``). Fall back to presence of
|
||||
``supportedInterfaces`` (a 1.0-only field) only when the explicit field is
|
||||
absent or unrecognized; cards carrying neither signal are treated as 0.3.
|
||||
"""
|
||||
pv = card.get("protocolVersion")
|
||||
if pv == "1.0":
|
||||
return "1.0"
|
||||
if pv == "0.3":
|
||||
return "0.3"
|
||||
# No protocolVersion field: use structural heuristic.
|
||||
normalized = normalize_protocol_version(card.get("protocolVersion"))
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
return "1.0" if "supportedInterfaces" in card else "0.3"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKey
|
|||
from litellm.proxy.a2a.agent_card import (
|
||||
SUPPORTED_A2A_PROTOCOL_VERSIONS,
|
||||
merge_agent_card,
|
||||
normalize_protocol_version,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
|
|
@ -51,7 +52,7 @@ def _proxy_base_url(http_request: Request) -> str:
|
|||
def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None:
|
||||
"""Reject an agent card pinning an unsupported A2A protocol version."""
|
||||
version = upstream_card.get("protocolVersion") if upstream_card else None
|
||||
if version is not None and version not in SUPPORTED_A2A_PROTOCOL_VERSIONS:
|
||||
if version is not None and normalize_protocol_version(version) is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import re
|
|||
import sys
|
||||
from functools import lru_cache
|
||||
from logging import Logger
|
||||
from typing import Any, Dict, FrozenSet, List, Mapping, Optional, Tuple, Union
|
||||
from typing import Any, Dict, FrozenSet, Iterator, List, Mapping, Optional, Tuple, Union
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
||||
|
|
@ -12,7 +12,12 @@ from litellm import Router, provider_list
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.litellm_core_utils.url_utils import (
|
||||
SSRFError,
|
||||
is_url_destination_allowed_by_host,
|
||||
provider_url_destination_candidates,
|
||||
validate_url,
|
||||
)
|
||||
from litellm.proxy._types import *
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
||||
|
|
@ -290,6 +295,7 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
|
|||
"use_ssl",
|
||||
# SDK-only field; also rejected outright in is_request_body_safe.
|
||||
"model_list",
|
||||
"vertex_ai_credentials",
|
||||
# Observability credentials, hosts, and project identifiers: derived
|
||||
# from the canonical ``_supported_callback_params`` allowlist so new
|
||||
# integrations are covered automatically. Sorted for stable iteration
|
||||
|
|
@ -342,6 +348,60 @@ def _check_banned_params(
|
|||
)
|
||||
|
||||
|
||||
_FALLBACK_FIELDS: tuple[str, ...] = (
|
||||
"fallbacks",
|
||||
"context_window_fallbacks",
|
||||
"content_policy_fallbacks",
|
||||
)
|
||||
|
||||
|
||||
def _iter_fallback_field_values(request_body: Mapping[str, object]) -> Iterator[object]:
|
||||
override = request_body.get("router_settings_override")
|
||||
for source in (request_body, override):
|
||||
if isinstance(source, Mapping):
|
||||
for field in _FALLBACK_FIELDS:
|
||||
yield source.get(field)
|
||||
|
||||
|
||||
def _iter_fallback_targets(value: object, depth: int) -> Iterator[str | Mapping[str, object]]:
|
||||
if depth > 2 * litellm.ROUTER_MAX_FALLBACKS:
|
||||
raise ValueError("Rejected Request: fallback nesting exceeds the allowed validation depth.")
|
||||
if not isinstance(value, list):
|
||||
return
|
||||
for item in value:
|
||||
if isinstance(item, str):
|
||||
yield item
|
||||
elif isinstance(item, Mapping):
|
||||
values = tuple(item.values())
|
||||
if not (values and all(isinstance(v, list) for v in values)):
|
||||
yield item
|
||||
if isinstance(item.get("model"), str):
|
||||
for field in _FALLBACK_FIELDS:
|
||||
yield from _iter_fallback_targets(item.get(field), depth + 1)
|
||||
else:
|
||||
for target_list in values:
|
||||
yield from _iter_fallback_targets(target_list, depth + 1)
|
||||
|
||||
|
||||
def iter_request_fallback_targets(request_body: Mapping[str, object]) -> Iterator[str | Mapping[str, object]]:
|
||||
for value in _iter_fallback_field_values(request_body):
|
||||
yield from _iter_fallback_targets(value, 0)
|
||||
|
||||
|
||||
def _reject_url_valued_fallback_target(value: str) -> None:
|
||||
allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
|
||||
for candidate in provider_url_destination_candidates(value):
|
||||
if not candidate.lower().startswith(("http://", "https://")):
|
||||
continue
|
||||
if is_url_destination_allowed_by_host(candidate, allowed_hosts):
|
||||
continue
|
||||
raise ValueError(
|
||||
f"Rejected Request: URL-valued fallback destination '{value}' is not allowed. "
|
||||
"Configure custom endpoints with api_base instead, or add the destination host to "
|
||||
"`provider_url_destination_allowed_hosts` in litellm_settings."
|
||||
)
|
||||
|
||||
|
||||
def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str) -> bool:
|
||||
"""
|
||||
Check if the request body is safe.
|
||||
|
|
@ -379,6 +439,14 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
|
|||
metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
|
||||
if metadata is not None:
|
||||
_check_banned_params(metadata, general_settings, llm_router, model)
|
||||
for target in iter_request_fallback_targets(request_body):
|
||||
if isinstance(target, dict):
|
||||
_check_banned_params(target, general_settings, llm_router, model)
|
||||
target_model = target.get("model")
|
||||
if isinstance(target_model, str):
|
||||
_reject_url_valued_fallback_target(target_model)
|
||||
elif isinstance(target, str):
|
||||
_reject_url_valued_fallback_target(target)
|
||||
litellm_params = _coerce_metadata_to_dict(request_body.get("litellm_params"))
|
||||
if litellm_params is not None:
|
||||
litellm_params_metadata = _coerce_metadata_to_dict(litellm_params.get("metadata"))
|
||||
|
|
|
|||
|
|
@ -1155,6 +1155,7 @@ class JWTAuthManager:
|
|||
org_id: Optional[str],
|
||||
api_key: str,
|
||||
jwt_valid_token: Optional[dict] = None,
|
||||
user_email: str | None = None,
|
||||
) -> Optional[JWTAuthBuilderResult]:
|
||||
"""Check admin status and route access permissions"""
|
||||
if not jwt_handler.is_admin(scopes=scopes):
|
||||
|
|
@ -1179,6 +1180,7 @@ class JWTAuthManager:
|
|||
token=api_key,
|
||||
team_id=None,
|
||||
user_id=user_id,
|
||||
user_email=user_email,
|
||||
end_user_id=None,
|
||||
org_id=org_id,
|
||||
team_membership=None,
|
||||
|
|
@ -2068,7 +2070,7 @@ class JWTAuthManager:
|
|||
|
||||
# Check admin access
|
||||
admin_result = await JWTAuthManager.check_admin_access(
|
||||
jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token
|
||||
jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token, user_email=user_email
|
||||
)
|
||||
if admin_result:
|
||||
await JWTAuthManager._attach_team_from_header_for_admin(
|
||||
|
|
@ -2303,6 +2305,7 @@ class JWTAuthManager:
|
|||
team_id=team_id,
|
||||
team_object=team_object,
|
||||
user_id=user_id,
|
||||
user_email=(user_object.user_email if user_object is not None and user_object.user_email else user_email),
|
||||
user_object=user_object,
|
||||
org_id=resolved_org_id, # Use resolved org_id (from alias lookup if applicable)
|
||||
org_object=org_object,
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ import secrets
|
|||
|
||||
import orjson
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Iterator, NamedTuple, List, Optional, Protocol, Tuple, Union, cast
|
||||
from typing import Any, Dict, NamedTuple, List, Optional, Protocol, Tuple, Union, cast
|
||||
|
||||
import fastapi
|
||||
from fastapi import HTTPException, Request, WebSocket, status
|
||||
|
|
@ -58,6 +58,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
get_model_from_request,
|
||||
get_request_route,
|
||||
get_request_route_template,
|
||||
iter_request_fallback_targets,
|
||||
normalize_request_route,
|
||||
pre_db_read_auth_checks,
|
||||
route_in_additonal_public_routes,
|
||||
|
|
@ -1011,7 +1012,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
|
|||
return
|
||||
if getattr(request.state, "parent_otel_span", None) is not None:
|
||||
return
|
||||
start_time = datetime.now()
|
||||
start_time = datetime.now(timezone.utc)
|
||||
try:
|
||||
request.state.litellm_received_at = start_time
|
||||
except Exception:
|
||||
|
|
@ -1061,7 +1062,7 @@ async def _user_api_key_auth_builder(
|
|||
# Prefer the receive-instant stamped by the early helper in
|
||||
# user_api_key_auth (before body parse) — overwriting it would shorten
|
||||
# the preprocessing-duration measurement by the body-parse window.
|
||||
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now()
|
||||
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now(timezone.utc)
|
||||
try:
|
||||
request.state.litellm_received_at = start_time
|
||||
except Exception:
|
||||
|
|
@ -1255,6 +1256,7 @@ async def _user_api_key_auth_builder(
|
|||
team_id = result["team_id"]
|
||||
team_object = result["team_object"]
|
||||
user_id = result["user_id"]
|
||||
user_email = result["user_email"]
|
||||
user_object = result["user_object"]
|
||||
end_user_id = result["end_user_id"]
|
||||
org_id = result["org_id"]
|
||||
|
|
@ -1279,6 +1281,7 @@ async def _user_api_key_auth_builder(
|
|||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id=user_id,
|
||||
user_email=user_email,
|
||||
team_id=team_id,
|
||||
team_alias=(team_object.team_alias if team_object is not None else None),
|
||||
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
|
||||
|
|
@ -1304,6 +1307,7 @@ async def _user_api_key_auth_builder(
|
|||
else LitellmUserRoles.INTERNAL_USER
|
||||
),
|
||||
user_id=user_id,
|
||||
user_email=user_email,
|
||||
org_id=org_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
end_user_id=end_user_id,
|
||||
|
|
@ -1345,6 +1349,7 @@ async def _user_api_key_auth_builder(
|
|||
)
|
||||
if auto_registered is not None:
|
||||
auto_registered.jwt_claims = jwt_claims
|
||||
auto_registered.user_email = user_email
|
||||
valid_token = auto_registered
|
||||
api_key = valid_token.token or ""
|
||||
|
||||
|
|
@ -2607,7 +2612,7 @@ async def _return_user_api_key_auth_obj(
|
|||
start_time: datetime,
|
||||
user_role: Optional[LitellmUserRoles] = None,
|
||||
) -> UserAPIKeyAuth:
|
||||
end_time = datetime.now()
|
||||
end_time = datetime.now(timezone.utc)
|
||||
|
||||
asyncio.create_task(
|
||||
user_api_key_service_logger_obj.async_service_success_hook(
|
||||
|
|
@ -2696,9 +2701,10 @@ def _update_key_budget_with_temp_budget_increase(
|
|||
) -> UserAPIKeyAuth:
|
||||
if valid_token.max_budget is None:
|
||||
return valid_token
|
||||
temp_budget_increase = _get_temp_budget_increase(valid_token) or 0.0
|
||||
valid_token.max_budget = valid_token.max_budget + temp_budget_increase
|
||||
return valid_token
|
||||
temp_budget_increase = _get_temp_budget_increase(valid_token)
|
||||
if not temp_budget_increase:
|
||||
return valid_token
|
||||
return valid_token.model_copy(update={"max_budget": valid_token.max_budget + temp_budget_increase})
|
||||
|
||||
|
||||
async def _lookup_end_user_and_apply_budget(
|
||||
|
|
@ -2796,19 +2802,11 @@ async def _enforce_key_and_fallback_model_access(
|
|||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Validate every fallback model name reachable by this request.
|
||||
# All three fields (``fallbacks``, ``context_window_fallbacks``,
|
||||
# ``content_policy_fallbacks``) are forwarded to the router as
|
||||
# per-request kwargs whether they appear at the top level of
|
||||
# ``request_data`` or nested under ``router_settings_override``.
|
||||
# Both surfaces must be validated against the API key's model
|
||||
# allowlist or a caller can smuggle a restricted model. VERIA-44.
|
||||
fallback_names: List[str] = []
|
||||
override_settings = request_data.get("router_settings_override")
|
||||
for _fb_key in ROUTER_FALLBACK_FIELDS:
|
||||
fallback_names.extend(iter_router_fallback_model_names(request_data.get(_fb_key)))
|
||||
if isinstance(override_settings, dict):
|
||||
fallback_names.extend(iter_router_fallback_model_names(override_settings.get(_fb_key)))
|
||||
fallback_names = tuple(
|
||||
name
|
||||
for target in iter_request_fallback_targets(request_data)
|
||||
if (name := _fallback_target_model_name(target)) is not None
|
||||
)
|
||||
|
||||
for _name in dict.fromkeys(fallback_names): # dedupe, preserve order
|
||||
await can_key_call_model(
|
||||
|
|
@ -2824,36 +2822,14 @@ async def _enforce_key_and_fallback_model_access(
|
|||
)
|
||||
|
||||
|
||||
ROUTER_FALLBACK_FIELDS: Tuple[str, ...] = (
|
||||
"fallbacks",
|
||||
"context_window_fallbacks",
|
||||
"content_policy_fallbacks",
|
||||
)
|
||||
|
||||
|
||||
def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]:
|
||||
"""Yield leaf model names from any of the supported fallbacks shapes.
|
||||
|
||||
Handles the simple top-level shape (``str`` or ``{"model": str}``) and
|
||||
the nested router-config shape (``[{primary: [fallback_list]}]``).
|
||||
"""
|
||||
if not isinstance(fallbacks, list):
|
||||
return
|
||||
for entry in fallbacks:
|
||||
if isinstance(entry, str):
|
||||
yield entry
|
||||
elif isinstance(entry, dict):
|
||||
if isinstance(entry.get("model"), str):
|
||||
yield entry["model"]
|
||||
continue
|
||||
for fallback_list in entry.values():
|
||||
if not isinstance(fallback_list, list):
|
||||
continue
|
||||
for m in fallback_list:
|
||||
if isinstance(m, str):
|
||||
yield m
|
||||
elif isinstance(m, dict) and isinstance(m.get("model"), str):
|
||||
yield m["model"]
|
||||
def _fallback_target_model_name(target: object) -> str | None:
|
||||
if isinstance(target, str):
|
||||
return target
|
||||
if isinstance(target, dict):
|
||||
model = target.get("model")
|
||||
if isinstance(model, str):
|
||||
return model
|
||||
return None
|
||||
|
||||
|
||||
async def _run_post_custom_auth_checks(
|
||||
|
|
|
|||
|
|
@ -523,7 +523,7 @@ export LITELLM_PROXY_API_KEY=sk-...
|
|||
lite model-groups list [--format table|json]
|
||||
```
|
||||
|
||||
Lists the model groups your key can reach on the proxy, via `/model_group/info`, along with each group's mode (`chat`, `embedding`, etc.) and per-token pricing. This is also what `lite autoroute configure` uses internally to discover what it can offer you.
|
||||
Lists the model groups your key can reach on the proxy, via `/model_group/info`, along with each group's mode (`chat`, `embedding`, etc.) and per-token pricing. Note this route needs management access; `lite autoroute configure` instead discovers models through `/v1/models`, so it works with a key scoped to just the AI API routes
|
||||
|
||||
#### Configure the Auto-Router
|
||||
|
||||
|
|
|
|||
|
|
@ -15,41 +15,25 @@ class DiscoveredModel(BaseModel):
|
|||
|
||||
name: str
|
||||
mode: str = "chat"
|
||||
input_cost_per_token: float | None = None
|
||||
output_cost_per_token: float | None = None
|
||||
|
||||
|
||||
class _RawModelGroup(BaseModel):
|
||||
class _RawModelListing(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
model_group: str
|
||||
# Optional: some real deployments return an explicit `"mode": null` for models that
|
||||
# were registered without a mode (seen for embedding models like voyage-4-large).
|
||||
# ModelGroupInfo's own "chat" default (litellm/types/router.py) only applies when the
|
||||
# key is missing entirely, not when it's present as null, so this must tolerate None.
|
||||
mode: str | None = "chat"
|
||||
input_cost_per_token: float | None = None
|
||||
output_cost_per_token: float | None = None
|
||||
id: str
|
||||
# /v1/models attaches "mode" (sourced from the cost map) only for models it can resolve;
|
||||
# a model whose mode is unknown arrives without the field, so default it to chat rather
|
||||
# than dropping it, which keeps it selectable as a routing target in the wizard.
|
||||
mode: str = "chat"
|
||||
|
||||
|
||||
_RAW_MODEL_GROUPS_ADAPTER = TypeAdapter(list[_RawModelGroup])
|
||||
_RAW_MODEL_LISTING_ADAPTER = TypeAdapter(list[_RawModelListing])
|
||||
|
||||
|
||||
def parse_discovered_models(raw: list[JsonValue]) -> tuple[DiscoveredModel, ...]:
|
||||
"""Validate a raw `/model_group/info` response into typed models."""
|
||||
parsed = _RAW_MODEL_GROUPS_ADAPTER.validate_python(raw)
|
||||
return tuple(
|
||||
DiscoveredModel(
|
||||
name=group.model_group,
|
||||
# A null mode means the server genuinely doesn't know what this model does;
|
||||
# "unknown" (rather than guessing "chat") keeps it out of both chat_models()
|
||||
# and embedding_models() instead of risking a wrong-mode deployment.
|
||||
mode=group.mode or "unknown",
|
||||
input_cost_per_token=group.input_cost_per_token,
|
||||
output_cost_per_token=group.output_cost_per_token,
|
||||
)
|
||||
for group in parsed
|
||||
)
|
||||
"""Validate a raw `/v1/models` response into typed models."""
|
||||
parsed = _RAW_MODEL_LISTING_ADAPTER.validate_python(raw)
|
||||
return tuple(DiscoveredModel(name=item.id, mode=item.mode) for item in parsed)
|
||||
|
||||
|
||||
def chat_models(models: tuple[DiscoveredModel, ...]) -> tuple[DiscoveredModel, ...]:
|
||||
|
|
|
|||
|
|
@ -111,12 +111,12 @@ def run_configure_wizard(ctx: click.Context) -> Path:
|
|||
api_key = ctx.obj["api_key"]
|
||||
client = Client(base_url=base_url, api_key=api_key)
|
||||
|
||||
raw_groups = client.model_groups.info()
|
||||
if not isinstance(raw_groups, list):
|
||||
raw_models = client.models.list()
|
||||
if not isinstance(raw_models, list):
|
||||
raise click.ClickException(
|
||||
f"Unexpected response from /model_group/info: expected a list, got {type(raw_groups).__name__}"
|
||||
f"Unexpected response from /v1/models: expected a list, got {type(raw_models).__name__}"
|
||||
)
|
||||
discovered = parse_discovered_models(raw_groups)
|
||||
discovered = parse_discovered_models(raw_models)
|
||||
chat_pool = chat_models(discovered)
|
||||
embedding_pool = embedding_models(discovered)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import copy
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional
|
||||
|
||||
import litellm
|
||||
|
|
@ -564,11 +565,8 @@ def process_callback(_callback: str, callback_type: str, environment_variables:
|
|||
|
||||
env_vars_dict: dict[str, str | None] = {}
|
||||
for _var in env_vars:
|
||||
env_variable = environment_variables.get(_var, None)
|
||||
if env_variable is None:
|
||||
env_vars_dict[_var] = None
|
||||
else:
|
||||
env_vars_dict[_var] = env_variable
|
||||
stored_value = environment_variables.get(_var, None)
|
||||
env_vars_dict[_var] = stored_value if stored_value is not None else os.getenv(_var)
|
||||
|
||||
return {"name": _callback, "variables": env_vars_dict, "type": callback_type}
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,11 @@ from litellm.proxy._types import (
|
|||
LiteLLM_UserTable,
|
||||
LiteLLM_VerificationToken,
|
||||
)
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
BudgetResetSettings,
|
||||
compute_budget_reset_at,
|
||||
get_budget_reset_settings,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
|
|
@ -32,9 +37,15 @@ class ResetBudgetJob:
|
|||
Resets the budget for all the keys, users, and teams that need it
|
||||
"""
|
||||
|
||||
def __init__(self, proxy_logging_obj: ProxyLogging, prisma_client: PrismaClient):
|
||||
def __init__(
|
||||
self,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prisma_client: PrismaClient,
|
||||
reset_settings: BudgetResetSettings | None = None,
|
||||
):
|
||||
self.proxy_logging_obj: ProxyLogging = proxy_logging_obj
|
||||
self.prisma_client: PrismaClient = prisma_client
|
||||
self.reset_settings: BudgetResetSettings = reset_settings or get_budget_reset_settings()
|
||||
|
||||
async def reset_budget(
|
||||
self,
|
||||
|
|
@ -237,7 +248,7 @@ class ResetBudgetJob:
|
|||
|
||||
if budgets_to_reset is not None and len(budgets_to_reset) > 0:
|
||||
for budget in budgets_to_reset:
|
||||
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now)
|
||||
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings)
|
||||
|
||||
await self.prisma_client.update_data(
|
||||
query_type="update_many",
|
||||
|
|
@ -442,7 +453,11 @@ class ResetBudgetJob:
|
|||
if keys_to_reset is not None and len(keys_to_reset) > 0:
|
||||
for key in keys_to_reset:
|
||||
try:
|
||||
updated_key = await ResetBudgetJob._reset_budget_for_key(key=key, current_time=now)
|
||||
updated_key = await ResetBudgetJob._reset_budget_for_key(
|
||||
key=key,
|
||||
current_time=now,
|
||||
reset_settings=self.reset_settings,
|
||||
)
|
||||
if updated_key is not None:
|
||||
updated_keys.append(updated_key)
|
||||
else:
|
||||
|
|
@ -513,7 +528,11 @@ class ResetBudgetJob:
|
|||
if users_to_reset is not None and len(users_to_reset) > 0:
|
||||
for user in users_to_reset:
|
||||
try:
|
||||
updated_user = await ResetBudgetJob._reset_budget_for_user(user=user, current_time=now)
|
||||
updated_user = await ResetBudgetJob._reset_budget_for_user(
|
||||
user=user,
|
||||
current_time=now,
|
||||
reset_settings=self.reset_settings,
|
||||
)
|
||||
if updated_user is not None:
|
||||
updated_users.append(updated_user)
|
||||
else:
|
||||
|
|
@ -588,7 +607,11 @@ class ResetBudgetJob:
|
|||
if teams_to_reset is not None and len(teams_to_reset) > 0:
|
||||
for team in teams_to_reset:
|
||||
try:
|
||||
updated_team = await ResetBudgetJob._reset_budget_for_team(team=team, current_time=now)
|
||||
updated_team = await ResetBudgetJob._reset_budget_for_team(
|
||||
team=team,
|
||||
current_time=now,
|
||||
reset_settings=self.reset_settings,
|
||||
)
|
||||
if updated_team is not None:
|
||||
updated_teams.append(updated_team)
|
||||
else:
|
||||
|
|
@ -655,10 +678,9 @@ class ResetBudgetJob:
|
|||
counter_key: str,
|
||||
spend_counter_cache: Any,
|
||||
now: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> bool:
|
||||
"""Reset a single budget window if expired. Returns True if the window was reset."""
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
|
||||
reset_at_str = window.get("reset_at")
|
||||
if not reset_at_str:
|
||||
return False
|
||||
|
|
@ -671,7 +693,9 @@ class ResetBudgetJob:
|
|||
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=0.0)
|
||||
except Exception as redis_err:
|
||||
verbose_proxy_logger.warning("Failed to reset Redis counter %s: %s", counter_key, redis_err)
|
||||
window["reset_at"] = get_budget_reset_time(budget_duration=window["budget_duration"]).isoformat()
|
||||
window["reset_at"] = compute_budget_reset_at(
|
||||
budget_duration=window["budget_duration"], settings=reset_settings
|
||||
).isoformat()
|
||||
return True
|
||||
|
||||
async def reset_budget_windows(self) -> None:
|
||||
|
|
@ -703,7 +727,13 @@ class ResetBudgetJob:
|
|||
changed = False
|
||||
for window in windows:
|
||||
counter_key = f"spend:key:{row['token']}:window:{window['budget_duration']}"
|
||||
if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now):
|
||||
if await ResetBudgetJob._reset_expired_window(
|
||||
window,
|
||||
counter_key,
|
||||
spend_counter_cache,
|
||||
now,
|
||||
self.reset_settings,
|
||||
):
|
||||
changed = True
|
||||
if changed:
|
||||
await VerificationTokenRepository(self.prisma_client).table.update(
|
||||
|
|
@ -726,7 +756,13 @@ class ResetBudgetJob:
|
|||
changed = False
|
||||
for window in windows:
|
||||
counter_key = f"spend:team:{row['team_id']}:window:{window['budget_duration']}"
|
||||
if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now):
|
||||
if await ResetBudgetJob._reset_expired_window(
|
||||
window,
|
||||
counter_key,
|
||||
spend_counter_cache,
|
||||
now,
|
||||
self.reset_settings,
|
||||
):
|
||||
changed = True
|
||||
if changed:
|
||||
await TeamRepository(self.prisma_client).table.update(
|
||||
|
|
@ -741,6 +777,7 @@ class ResetBudgetJob:
|
|||
item: Union[LiteLLM_TeamTable, LiteLLM_UserTable, LiteLLM_VerificationToken],
|
||||
current_time: datetime,
|
||||
item_type: Literal["key", "team", "user"],
|
||||
reset_settings: BudgetResetSettings,
|
||||
):
|
||||
"""
|
||||
In-place, updates spend=0, and sets budget_reset_at to current_time + budget_duration
|
||||
|
|
@ -755,24 +792,40 @@ class ResetBudgetJob:
|
|||
try:
|
||||
item.spend = 0.0
|
||||
if hasattr(item, "budget_duration") and item.budget_duration is not None:
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
get_budget_reset_time,
|
||||
item.budget_reset_at = compute_budget_reset_at(
|
||||
budget_duration=item.budget_duration, settings=reset_settings
|
||||
)
|
||||
|
||||
item.budget_reset_at = get_budget_reset_time(budget_duration=item.budget_duration)
|
||||
return item
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error resetting budget for %s: %s. Item: %s", item_type, e, item)
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
async def _reset_budget_for_team(team: LiteLLM_TeamTable, current_time: datetime) -> Optional[LiteLLM_TeamTable]:
|
||||
await ResetBudgetJob._reset_budget_common(item=team, current_time=current_time, item_type="team")
|
||||
async def _reset_budget_for_team(
|
||||
team: LiteLLM_TeamTable,
|
||||
current_time: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> LiteLLM_TeamTable | None:
|
||||
await ResetBudgetJob._reset_budget_common(
|
||||
item=team,
|
||||
current_time=current_time,
|
||||
item_type="team",
|
||||
reset_settings=reset_settings,
|
||||
)
|
||||
return team
|
||||
|
||||
@staticmethod
|
||||
async def _reset_budget_for_user(user: LiteLLM_UserTable, current_time: datetime) -> Optional[LiteLLM_UserTable]:
|
||||
await ResetBudgetJob._reset_budget_common(item=user, current_time=current_time, item_type="user")
|
||||
async def _reset_budget_for_user(
|
||||
user: LiteLLM_UserTable,
|
||||
current_time: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> LiteLLM_UserTable | None:
|
||||
await ResetBudgetJob._reset_budget_common(
|
||||
item=user,
|
||||
current_time=current_time,
|
||||
item_type="user",
|
||||
reset_settings=reset_settings,
|
||||
)
|
||||
return user
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -788,15 +841,15 @@ class ResetBudgetJob:
|
|||
|
||||
@staticmethod
|
||||
async def _reset_budget_reset_at_date(
|
||||
budget: LiteLLM_BudgetTableFull, current_time: datetime
|
||||
budget: LiteLLM_BudgetTableFull,
|
||||
current_time: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> LiteLLM_BudgetTableFull:
|
||||
try:
|
||||
if budget.budget_duration is not None:
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
get_budget_reset_time,
|
||||
budget.budget_reset_at = compute_budget_reset_at(
|
||||
budget_duration=budget.budget_duration, settings=reset_settings
|
||||
)
|
||||
|
||||
budget.budget_reset_at = get_budget_reset_time(budget_duration=budget.budget_duration)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget)
|
||||
raise e
|
||||
|
|
@ -804,7 +857,14 @@ class ResetBudgetJob:
|
|||
|
||||
@staticmethod
|
||||
async def _reset_budget_for_key(
|
||||
key: LiteLLM_VerificationToken, current_time: datetime
|
||||
) -> Optional[LiteLLM_VerificationToken]:
|
||||
await ResetBudgetJob._reset_budget_common(item=key, current_time=current_time, item_type="key")
|
||||
key: LiteLLM_VerificationToken,
|
||||
current_time: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> LiteLLM_VerificationToken | None:
|
||||
await ResetBudgetJob._reset_budget_common(
|
||||
item=key,
|
||||
current_time=current_time,
|
||||
item_type="key",
|
||||
reset_settings=reset_settings,
|
||||
)
|
||||
return key
|
||||
|
|
|
|||
|
|
@ -1,10 +1,47 @@
|
|||
from datetime import datetime, timezone
|
||||
from datetime import datetime, time, timezone
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
|
||||
|
||||
|
||||
def get_budget_reset_timezone():
|
||||
class BudgetResetSettings(BaseModel):
|
||||
"""Immutable, validated settings that govern when budgets reset.
|
||||
|
||||
Parsed once from `litellm_settings` and injected into consumers (the reset
|
||||
job, management endpoints) so reset times never depend on reaching into
|
||||
module-level globals at call time.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
timezone: str = "UTC"
|
||||
reset_time_of_day: time = time(0, 0)
|
||||
|
||||
|
||||
def parse_budget_reset_time(raw: object) -> time:
|
||||
"""Parse a `budget_reset_time` config value (e.g. "12:00") into a `time`.
|
||||
|
||||
Falls back to midnight when unset; raises a clear error on a malformed value
|
||||
so a bad config fails loudly at startup instead of silently resetting at midnight.
|
||||
"""
|
||||
if raw is None or raw == "":
|
||||
return time(0, 0)
|
||||
if not isinstance(raw, str):
|
||||
raise ValueError(f"Invalid budget_reset_time {raw!r}; must be a quoted 24-hour 'HH:MM' string, e.g. \"12:00\"")
|
||||
for fmt in ("%H:%M", "%H:%M:%S"):
|
||||
try:
|
||||
parsed = datetime.strptime(raw, fmt)
|
||||
return time(hour=parsed.hour, minute=parsed.minute, second=parsed.second)
|
||||
except ValueError:
|
||||
continue
|
||||
raise ValueError(
|
||||
f"Invalid budget_reset_time {raw!r}; expected a 24-hour 'HH:MM' or 'HH:MM:SS' string, e.g. \"12:00\""
|
||||
)
|
||||
|
||||
|
||||
def get_budget_reset_timezone() -> str:
|
||||
"""
|
||||
Get the budget reset timezone from litellm_settings.
|
||||
Falls back to UTC if not specified.
|
||||
|
|
@ -15,15 +52,29 @@ def get_budget_reset_timezone():
|
|||
return getattr(litellm, "timezone", None) or "UTC"
|
||||
|
||||
|
||||
def get_budget_reset_time(budget_duration: str) -> datetime:
|
||||
"""
|
||||
Get the budget reset time based on the configured timezone.
|
||||
Falls back to UTC if not specified.
|
||||
"""
|
||||
def get_budget_reset_settings() -> BudgetResetSettings:
|
||||
"""Build validated reset settings from litellm_settings. Raises on a malformed
|
||||
`budget_reset_time`, which lets the proxy fail fast at startup."""
|
||||
return BudgetResetSettings(
|
||||
timezone=get_budget_reset_timezone(),
|
||||
reset_time_of_day=parse_budget_reset_time(getattr(litellm, "budget_reset_time", None)),
|
||||
)
|
||||
|
||||
reset_at = get_next_standardized_reset_time(
|
||||
|
||||
def compute_budget_reset_at(budget_duration: str, settings: BudgetResetSettings) -> datetime:
|
||||
"""Compute the next reset time for a budget duration using injected settings."""
|
||||
return get_next_standardized_reset_time(
|
||||
duration=budget_duration,
|
||||
current_time=datetime.now(timezone.utc),
|
||||
timezone_str=get_budget_reset_timezone(),
|
||||
timezone_str=settings.timezone,
|
||||
reset_time_of_day=settings.reset_time_of_day,
|
||||
)
|
||||
return reset_at
|
||||
|
||||
|
||||
def get_budget_reset_time(budget_duration: str) -> datetime:
|
||||
"""Get the budget reset time using the globally-configured timezone and reset time.
|
||||
|
||||
Thin wrapper over `compute_budget_reset_at` for callers that don't yet receive
|
||||
`BudgetResetSettings` by injection (creation/update endpoints, startup backfill).
|
||||
"""
|
||||
return compute_budget_reset_at(budget_duration, get_budget_reset_settings())
|
||||
|
|
|
|||
9
litellm/proxy/config_resolvers/__init__.py
Normal file
9
litellm/proxy/config_resolvers/__init__.py
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
"""Typed, provenance-aware resolution of proxy settings from DB then env."""
|
||||
|
||||
from litellm.proxy.config_resolvers._descriptors import (
|
||||
FieldDescriptor,
|
||||
FieldSource,
|
||||
resolve_fields,
|
||||
)
|
||||
|
||||
__all__ = ["FieldDescriptor", "FieldSource", "resolve_fields"]
|
||||
73
litellm/proxy/config_resolvers/_descriptors.py
Normal file
73
litellm/proxy/config_resolvers/_descriptors.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
"""Shared primitive for resolving a settings value from its sources.
|
||||
|
||||
A ``FieldDescriptor`` names, for one setting, where it lives in the stored DB
|
||||
row (``db_key``), which process env var carries it (``env_var``), whether it is
|
||||
a secret, and its effective default. ``resolve_fields`` reconciles a set of
|
||||
descriptors against a decrypted DB row and the process environment with a fixed
|
||||
precedence, returning the resolved values plus per-field provenance so a caller
|
||||
can tell whether a value came from the database, the environment, a default, or
|
||||
is unset.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
FieldSource = Literal["db", "env", "default", "unset"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FieldDescriptor:
|
||||
field_name: str
|
||||
db_key: str
|
||||
env_var: str
|
||||
is_secret: bool = False
|
||||
default: str | None = None
|
||||
|
||||
|
||||
def _db_is_set(db_value: object, empty_db_is_set: bool) -> bool:
|
||||
if empty_db_is_set:
|
||||
# A stored key that is present, even as "", is an explicit admin choice
|
||||
# (e.g. clearing an alerting webhook) and must win over a stale env var.
|
||||
return db_value is not None
|
||||
# A blank stored value is treated as absent, so it falls through to env. This
|
||||
# fits settings whose clear path also unsets the env var (e.g. SSO).
|
||||
return isinstance(db_value, str) and bool(db_value.strip())
|
||||
|
||||
|
||||
def _resolve_one(
|
||||
descriptor: FieldDescriptor,
|
||||
db_values: Mapping[str, object],
|
||||
env: Mapping[str, str],
|
||||
empty_db_is_set: bool,
|
||||
) -> tuple[str, str | None, FieldSource]:
|
||||
db_value = db_values.get(descriptor.db_key)
|
||||
if _db_is_set(db_value, empty_db_is_set):
|
||||
return descriptor.field_name, db_value if isinstance(db_value, str) else str(db_value), "db"
|
||||
env_value = env.get(descriptor.env_var)
|
||||
if isinstance(env_value, str) and env_value.strip():
|
||||
return descriptor.field_name, env_value, "env"
|
||||
if descriptor.default is not None:
|
||||
return descriptor.field_name, descriptor.default, "default"
|
||||
return descriptor.field_name, None, "unset"
|
||||
|
||||
|
||||
def resolve_fields(
|
||||
descriptors: Sequence[FieldDescriptor],
|
||||
db_values: Mapping[str, object],
|
||||
env: Mapping[str, str],
|
||||
empty_db_is_set: bool = False,
|
||||
) -> tuple[dict[str, str | None], dict[str, FieldSource]]:
|
||||
"""Resolve every descriptor to (values, provenance).
|
||||
|
||||
Precedence per field: a set stored value wins, else a non-blank process env
|
||||
var, else the descriptor default, else unset. ``empty_db_is_set`` selects
|
||||
how a present-but-empty stored value is read: ``False`` treats it as absent
|
||||
so it falls back to env (SSO, whose clear path also unsets the env var);
|
||||
``True`` treats it as an explicit clear that wins over env (alerting, whose
|
||||
clear path stores "" without unsetting the env var).
|
||||
"""
|
||||
resolved = tuple(_resolve_one(descriptor, db_values, env, empty_db_is_set) for descriptor in descriptors)
|
||||
values = {field_name: value for field_name, value, _ in resolved}
|
||||
provenance = {field_name: source for field_name, _, source in resolved}
|
||||
return values, provenance
|
||||
25
litellm/proxy/config_resolvers/alerting.py
Normal file
25
litellm/proxy/config_resolvers/alerting.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
"""Descriptor tables for the alerting settings surfaced by /get/config/callbacks.
|
||||
|
||||
These reconcile the stored ``environment_variables`` blob (keyed by the
|
||||
uppercase env-var names) with the process environment. SMTP_PORT and SMTP_TLS
|
||||
carry the same effective defaults the mail-send path applies, so the settings
|
||||
page shows the config that mail would actually use rather than a blank.
|
||||
"""
|
||||
|
||||
from litellm.proxy.config_resolvers._descriptors import FieldDescriptor
|
||||
|
||||
EMAIL_DESCRIPTORS: tuple[FieldDescriptor, ...] = (
|
||||
FieldDescriptor("SMTP_HOST", "SMTP_HOST", "SMTP_HOST"),
|
||||
FieldDescriptor("SMTP_PORT", "SMTP_PORT", "SMTP_PORT", default="587"),
|
||||
FieldDescriptor("SMTP_TLS", "SMTP_TLS", "SMTP_TLS", default="True"),
|
||||
FieldDescriptor("SMTP_USERNAME", "SMTP_USERNAME", "SMTP_USERNAME", is_secret=True),
|
||||
FieldDescriptor("SMTP_PASSWORD", "SMTP_PASSWORD", "SMTP_PASSWORD", is_secret=True),
|
||||
FieldDescriptor("SMTP_SENDER_EMAIL", "SMTP_SENDER_EMAIL", "SMTP_SENDER_EMAIL"),
|
||||
FieldDescriptor("TEST_EMAIL_ADDRESS", "TEST_EMAIL_ADDRESS", "TEST_EMAIL_ADDRESS"),
|
||||
FieldDescriptor("EMAIL_LOGO_URL", "EMAIL_LOGO_URL", "EMAIL_LOGO_URL"),
|
||||
FieldDescriptor("EMAIL_SUPPORT_CONTACT", "EMAIL_SUPPORT_CONTACT", "EMAIL_SUPPORT_CONTACT"),
|
||||
)
|
||||
|
||||
SLACK_DESCRIPTORS: tuple[FieldDescriptor, ...] = (
|
||||
FieldDescriptor("SLACK_WEBHOOK_URL", "SLACK_WEBHOOK_URL", "SLACK_WEBHOOK_URL", is_secret=True),
|
||||
)
|
||||
94
litellm/proxy/config_resolvers/sso.py
Normal file
94
litellm/proxy/config_resolvers/sso.py
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
"""Resolved SSO config object.
|
||||
|
||||
Reconciles the dedicated ``sso_config`` DB row (lowercase, per-value encrypted
|
||||
keys) with the process environment (uppercase env vars) into a typed
|
||||
``SSOConfig`` plus per-field provenance. This is the single source of truth for
|
||||
the SSO field -> env-var mapping, used by both the read-back endpoint and the
|
||||
save endpoint so the two can never drift.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.config_resolvers._descriptors import (
|
||||
FieldDescriptor,
|
||||
FieldSource,
|
||||
resolve_fields,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.ui_sso import (
|
||||
RoleMappings,
|
||||
SSOConfig,
|
||||
TeamMappings,
|
||||
)
|
||||
|
||||
SSO_DESCRIPTORS: tuple[FieldDescriptor, ...] = (
|
||||
FieldDescriptor("google_client_id", "google_client_id", "GOOGLE_CLIENT_ID"),
|
||||
FieldDescriptor("google_client_secret", "google_client_secret", "GOOGLE_CLIENT_SECRET", is_secret=True),
|
||||
FieldDescriptor("microsoft_client_id", "microsoft_client_id", "MICROSOFT_CLIENT_ID"),
|
||||
FieldDescriptor("microsoft_client_secret", "microsoft_client_secret", "MICROSOFT_CLIENT_SECRET", is_secret=True),
|
||||
FieldDescriptor("microsoft_tenant", "microsoft_tenant", "MICROSOFT_TENANT"),
|
||||
FieldDescriptor("generic_client_id", "generic_client_id", "GENERIC_CLIENT_ID"),
|
||||
FieldDescriptor("generic_client_secret", "generic_client_secret", "GENERIC_CLIENT_SECRET", is_secret=True),
|
||||
FieldDescriptor(
|
||||
"generic_authorization_endpoint", "generic_authorization_endpoint", "GENERIC_AUTHORIZATION_ENDPOINT"
|
||||
),
|
||||
FieldDescriptor("generic_token_endpoint", "generic_token_endpoint", "GENERIC_TOKEN_ENDPOINT"),
|
||||
FieldDescriptor("generic_userinfo_endpoint", "generic_userinfo_endpoint", "GENERIC_USERINFO_ENDPOINT"),
|
||||
FieldDescriptor("generic_scope", "generic_scope", "GENERIC_SCOPE", default="openid email profile"),
|
||||
FieldDescriptor("proxy_base_url", "proxy_base_url", "PROXY_BASE_URL"),
|
||||
)
|
||||
|
||||
# Derived from the descriptor table so read (masking) and the field->env mapping
|
||||
# never diverge from the resolver.
|
||||
SSO_SECRET_FIELDS: frozenset[str] = frozenset(d.field_name for d in SSO_DESCRIPTORS if d.is_secret)
|
||||
SSO_FIELD_ENV_VARS: dict[str, str] = {d.field_name: d.env_var for d in SSO_DESCRIPTORS}
|
||||
|
||||
# Structured sub-objects stored on the SSO row that are not simple env-backed
|
||||
# scalars; handled outside the descriptor resolution.
|
||||
_STRUCTURED_KEYS = ("role_mappings", "team_mappings")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedSSOConfig:
|
||||
config: SSOConfig
|
||||
provenance: dict[str, FieldSource]
|
||||
|
||||
|
||||
def _decrypt(raw: Mapping[str, object]) -> dict[str, object]:
|
||||
return {
|
||||
key: (
|
||||
decrypt_value_helper(value=value, key=key, return_original_value=True) if isinstance(value, str) else value
|
||||
)
|
||||
for key, value in raw.items()
|
||||
}
|
||||
|
||||
|
||||
def _parse_role_mappings(data: object) -> RoleMappings | None:
|
||||
# The stored row is JSON, so mappings arrive as a dict (or are absent).
|
||||
return RoleMappings(**data) if isinstance(data, dict) else None
|
||||
|
||||
|
||||
def _parse_team_mappings(data: object) -> TeamMappings | None:
|
||||
return TeamMappings(**data) if isinstance(data, dict) else None
|
||||
|
||||
|
||||
def resolve_sso_config(sso_db_settings: Mapping[str, object] | None, env: Mapping[str, str]) -> ResolvedSSOConfig:
|
||||
"""Resolve the effective SSO config: stored row first, then process env.
|
||||
|
||||
Decryption happens here, once, via the pure ``decrypt_value_helper``; this
|
||||
function never writes ``os.environ`` (unlike the legacy read path). Values
|
||||
are returned unmasked so the login path could consume them; the read-back
|
||||
endpoint is responsible for masking secrets before responding to the UI.
|
||||
"""
|
||||
raw = dict(sso_db_settings) if sso_db_settings else {}
|
||||
decrypted = _decrypt({key: value for key, value in raw.items() if key not in _STRUCTURED_KEYS})
|
||||
values, provenance = resolve_fields(SSO_DESCRIPTORS, decrypted, env)
|
||||
structured = {
|
||||
"user_email": decrypted.get("user_email"),
|
||||
"ui_access_mode": decrypted.get("ui_access_mode"),
|
||||
"role_mappings": _parse_role_mappings(raw.get("role_mappings")),
|
||||
"team_mappings": _parse_team_mappings(raw.get("team_mappings")),
|
||||
}
|
||||
config = SSOConfig(**{**values, **structured})
|
||||
return ResolvedSSOConfig(config=config, provenance=provenance)
|
||||
|
|
@ -2046,6 +2046,90 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
masking_index += 1
|
||||
verbose_proxy_logger.debug("Applied masking to choice text content")
|
||||
|
||||
@staticmethod
|
||||
def _incremental_scan_cache() -> DualCache:
|
||||
"""Resolve the cache used to remember which segments a session already scanned.
|
||||
|
||||
Prefers the proxy's shared cache (``internal_usage_cache.dual_cache``), which is
|
||||
backed by Redis when the deployment configures it, so incremental state is shared
|
||||
across proxy instances. Falls back to a process-local ``DualCache`` singleton when
|
||||
the proxy is not running (e.g. unit tests), where sharing does not apply.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import dc as fallback_cache
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging
|
||||
except Exception: # noqa: BLE001 # proxy not importable outside the server; use local fallback
|
||||
return fallback_cache
|
||||
if _proxy_logging is not None:
|
||||
return _proxy_logging.internal_usage_cache.dual_cache
|
||||
return fallback_cache
|
||||
|
||||
def _bedrock_response_has_masked_output(self, response: BedrockGuardrailResponse) -> bool:
|
||||
"""Return True if the guardrail rewrote (masked/anonymized) any scanned text.
|
||||
|
||||
Bedrock returns non-empty ``output``/``outputs`` text only when it changed the
|
||||
content; an ``action == "NONE"`` response leaves both empty.
|
||||
"""
|
||||
for field in ("output", "outputs"):
|
||||
items = response.get(field) or []
|
||||
if any(isinstance(item, dict) and item.get("text") for item in items):
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _apply_incremental_request_scan(
|
||||
self,
|
||||
texts: list[str],
|
||||
inputs: "GenericGuardrailAPIInputs",
|
||||
request_data: dict,
|
||||
) -> Optional["GenericGuardrailAPIInputs"]:
|
||||
"""Scan only the text segments not already seen earlier in this session.
|
||||
|
||||
Returns ``None`` when incremental scanning is inactive (feature off, no
|
||||
session id, masking enabled, or cache unavailable) or when the guardrail
|
||||
turns out to mask content, telling the caller to run the normal full scan.
|
||||
Otherwise scans only the new segments and skips the Bedrock call entirely
|
||||
when nothing is new. Incremental mode is for blocking/detection guardrails
|
||||
only: if the guardrail returns masked output it cannot be applied to the
|
||||
skipped context, so the scan falls back to the full path and no session
|
||||
state is recorded.
|
||||
"""
|
||||
cache = self._incremental_scan_cache()
|
||||
|
||||
new_texts = await self.filter_new_texts_for_session(
|
||||
texts=texts,
|
||||
request_data=request_data,
|
||||
cache=cache,
|
||||
)
|
||||
if new_texts is None:
|
||||
return None
|
||||
|
||||
if not new_texts:
|
||||
verbose_proxy_logger.debug("Bedrock Guardrail: no new messages to scan for this session, skipping API call")
|
||||
return inputs
|
||||
|
||||
bedrock_response = await self.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=[ChatCompletionUserMessage(role="user", content=text) for text in new_texts],
|
||||
request_data=request_data,
|
||||
logging_event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
|
||||
if self._bedrock_response_has_masked_output(bedrock_response):
|
||||
verbose_proxy_logger.warning(
|
||||
"Bedrock Guardrail %s: guardrail returned masked/anonymized content; "
|
||||
"only_scan_new_messages cannot apply masking to skipped context, falling back to a full-context scan",
|
||||
self.guardrail_name,
|
||||
)
|
||||
return None
|
||||
|
||||
await self.mark_texts_scanned(
|
||||
texts=texts,
|
||||
request_data=request_data,
|
||||
cache=cache,
|
||||
)
|
||||
return inputs
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: "GenericGuardrailAPIInputs",
|
||||
|
|
@ -2077,6 +2161,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
try:
|
||||
verbose_proxy_logger.debug(f"Bedrock Guardrail: Applying guardrail to {len(texts)} text(s)")
|
||||
|
||||
if input_type == "request":
|
||||
incremental_result = await self._apply_incremental_request_scan(
|
||||
texts=texts,
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
)
|
||||
if incremental_result is not None:
|
||||
return incremental_result
|
||||
|
||||
masked_texts = []
|
||||
|
||||
selection = self._select_messages_for_apply_guardrail(
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
mask_response_content=litellm_params.mask_response_content,
|
||||
fail_on_error=litellm_params.fail_on_error,
|
||||
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
|
||||
sanitize_error_detail=litellm_params.sanitize_error_detail,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from typing import (
|
|||
Union,
|
||||
)
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -35,7 +36,8 @@ from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
|
|||
MODEL_ARMOR_MAX_FILE_SIZE_BYTES,
|
||||
plan_file_scans,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
|
|
@ -50,6 +52,33 @@ from litellm.types.utils import (
|
|||
GUARDRAIL_NAME = "model_armor"
|
||||
|
||||
|
||||
class ModelArmorAPIError(Exception):
|
||||
"""Model Armor API failure (non-2xx), distinct from a content-block decision so
|
||||
hooks can honor fail_on_error. The detail is already sanitized per configuration."""
|
||||
|
||||
def __init__(self, detail: str):
|
||||
super().__init__(detail)
|
||||
self.detail = detail
|
||||
|
||||
|
||||
_SCANNED_CONTENT_KEYS = frozenset({"text", "sanitizedText", "findings", "maliciousUriMatchedItems"})
|
||||
|
||||
RedactablePayload = Union[dict, list, str, int, float, bool, None]
|
||||
|
||||
|
||||
def _redact_scanned_content(payload: RedactablePayload, depth: int = 0) -> RedactablePayload:
|
||||
if depth >= DEFAULT_MAX_RECURSE_DEPTH:
|
||||
return "[REDACTED]"
|
||||
if isinstance(payload, dict):
|
||||
return {
|
||||
key: "[REDACTED]" if key in _SCANNED_CONTENT_KEYS else _redact_scanned_content(value, depth + 1)
|
||||
for key, value in payload.items()
|
||||
}
|
||||
if isinstance(payload, list):
|
||||
return [_redact_scanned_content(item, depth + 1) for item in payload]
|
||||
return payload
|
||||
|
||||
|
||||
class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
||||
"""
|
||||
Google Cloud Model Armor Guardrail integration for LiteLLM.
|
||||
|
|
@ -76,6 +105,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
location: Optional[str] = None,
|
||||
credentials: Optional[Any] = None,
|
||||
api_endpoint: Optional[str] = None,
|
||||
sanitize_error_detail: "bool | None" = True,
|
||||
**kwargs,
|
||||
):
|
||||
# Set supported event hooks if not already provided
|
||||
|
|
@ -98,6 +128,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
self.location = location or "us-central1"
|
||||
self.credentials = credentials
|
||||
self.api_endpoint = api_endpoint
|
||||
self.sanitize_error_detail = sanitize_error_detail is not False
|
||||
|
||||
# Store optional params
|
||||
self.optional_params = kwargs
|
||||
|
|
@ -141,6 +172,67 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
verbose_proxy_logger.debug("Model Armor: Skipping non-ModelResponse type: %s", type(response).__name__)
|
||||
return ""
|
||||
|
||||
def _build_api_error_detail(self, status_code: int, response_text: str) -> str:
|
||||
if self.sanitize_error_detail:
|
||||
return f"Model Armor API error (upstream {status_code})"
|
||||
return f"Model Armor API error (upstream {status_code}): {response_text}"
|
||||
|
||||
def _build_block_error_detail(self, message: str, armor_response: RedactablePayload) -> dict:
|
||||
if self.sanitize_error_detail:
|
||||
return {"error": message}
|
||||
return {"error": message, "model_armor_response": armor_response}
|
||||
|
||||
def _build_logging_response(self, armor_response: RedactablePayload) -> RedactablePayload:
|
||||
if self.sanitize_error_detail:
|
||||
return _redact_scanned_content(armor_response)
|
||||
return armor_response
|
||||
|
||||
def _raise_if_fail_closed(self, e: ModelArmorAPIError) -> None:
|
||||
if self.optional_params.get("fail_on_error", True):
|
||||
raise e from None
|
||||
|
||||
def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
|
||||
super().update_in_memory_litellm_params(litellm_params)
|
||||
self.sanitize_error_detail = self.sanitize_error_detail is not False
|
||||
|
||||
def _log_request_debug(
|
||||
self,
|
||||
url: str,
|
||||
body: dict,
|
||||
file_bytes: "bytes | None",
|
||||
file_type: "str | None",
|
||||
) -> None:
|
||||
# Never log byteData: it is the full base64 of the scanned document. Log only its
|
||||
# type and size so debug deployments cannot leak the contents the guardrail inspects.
|
||||
if file_bytes is not None and file_type is not None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor file request - URL: %s, byteDataType: %s, bytes: %d",
|
||||
url,
|
||||
file_type,
|
||||
len(file_bytes),
|
||||
)
|
||||
elif self.sanitize_error_detail:
|
||||
verbose_proxy_logger.debug("Model Armor request - URL: %s", url)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor request - URL: %s, Body: %s",
|
||||
url,
|
||||
body,
|
||||
)
|
||||
|
||||
def _log_response_debug(self, status_code: int, response_text: str) -> None:
|
||||
if self.sanitize_error_detail:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor response - Status: %s",
|
||||
status_code,
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor response - Status: %s, Body: %s",
|
||||
status_code,
|
||||
response_text,
|
||||
)
|
||||
|
||||
async def make_model_armor_request(
|
||||
self,
|
||||
content: Optional[str] = None,
|
||||
|
|
@ -185,48 +277,37 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
"Authorization": f"Bearer {access_token}",
|
||||
}
|
||||
|
||||
# Never log byteData: it is the full base64 of the scanned document. Log only its
|
||||
# type and size so debug deployments cannot leak the contents the guardrail inspects.
|
||||
if file_bytes is not None and file_type is not None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor file request - URL: %s, byteDataType: %s, bytes: %d",
|
||||
url,
|
||||
file_type,
|
||||
len(file_bytes),
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor request - URL: %s, Body: %s",
|
||||
url,
|
||||
body,
|
||||
)
|
||||
self._log_request_debug(url=url, body=body, file_bytes=file_bytes, file_type=file_type)
|
||||
|
||||
# Make request
|
||||
if self.async_handler is None:
|
||||
raise ValueError("Async handler not initialized")
|
||||
|
||||
response = await self.async_handler.post(
|
||||
url=url,
|
||||
json=body,
|
||||
headers=headers,
|
||||
)
|
||||
try:
|
||||
response = await self.async_handler.post(
|
||||
url=url,
|
||||
json=body,
|
||||
headers=headers,
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
detail = self._build_api_error_detail(e.response.status_code, e.response.text)
|
||||
verbose_proxy_logger.error(
|
||||
"Model Armor API error - Status: %s, Detail: %s",
|
||||
e.response.status_code,
|
||||
detail,
|
||||
)
|
||||
raise ModelArmorAPIError(detail) from None
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Model Armor response - Status: %s, Body: %s",
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
self._log_response_debug(status_code=response.status_code, response_text=response.text)
|
||||
|
||||
if response.status_code != 200:
|
||||
detail = self._build_api_error_detail(response.status_code, response.text)
|
||||
verbose_proxy_logger.error(
|
||||
"Model Armor API error - Status: %s, Response: %s",
|
||||
"Model Armor API error - Status: %s, Detail: %s",
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Model Armor API error (upstream {response.status_code}): {response.text}",
|
||||
detail,
|
||||
)
|
||||
raise ModelArmorAPIError(detail)
|
||||
|
||||
json_response = response.json()
|
||||
if hasattr(json_response, "__await__"):
|
||||
|
|
@ -351,9 +432,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
Override to store only the Model Armor API response, not the entire data dict.
|
||||
This prevents circular references in logging.
|
||||
"""
|
||||
# Retrieve the Model Armor response & status stored on the per-request `metadata` object.
|
||||
metadata = request_data.get("metadata", {}) if isinstance(request_data, dict) else {}
|
||||
|
||||
guardrail_response = metadata.get("_model_armor_response", {})
|
||||
|
||||
# Determine status – default to "success" but prefer the explicit value if present.
|
||||
|
|
@ -444,6 +523,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
file_bytes=attachment.file_bytes,
|
||||
file_type=attachment.byte_data_type,
|
||||
)
|
||||
except ModelArmorAPIError as e:
|
||||
self._raise_if_fail_closed(e)
|
||||
continue
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -459,7 +541,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
# otherwise a PII-only (SDP deidentify) document would pass through unscrubbed.
|
||||
blocked = self._should_block_content(armor_response, allow_sanitization=False)
|
||||
metadata["_model_armor_response"] = self._append_armor_response(
|
||||
metadata.get("_model_armor_response"), armor_response
|
||||
metadata.get("_model_armor_response"),
|
||||
self._build_logging_response(armor_response),
|
||||
)
|
||||
if blocked or metadata.get("_model_armor_status") == "blocked":
|
||||
metadata["_model_armor_status"] = "blocked"
|
||||
|
|
@ -469,10 +552,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if blocked:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Content blocked by Model Armor",
|
||||
"model_armor_response": armor_response,
|
||||
},
|
||||
detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response),
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
|
|
@ -530,7 +610,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request
|
||||
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
|
||||
metadata["_model_armor_response"] = self._append_armor_response(
|
||||
metadata.get("_model_armor_response"), armor_response
|
||||
metadata.get("_model_armor_response"),
|
||||
self._build_logging_response(armor_response),
|
||||
)
|
||||
# Pre-compute guardrail status for downstream logging. A blocked response will eventually raise
|
||||
# an HTTPException, however in scenarios where the caller decides to ignore the exception (e.g.
|
||||
|
|
@ -548,10 +629,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if blocked:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Content blocked by Model Armor",
|
||||
"model_armor_response": armor_response,
|
||||
},
|
||||
detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response),
|
||||
)
|
||||
|
||||
# If mask_request_content is enabled, update messages with sanitized content
|
||||
|
|
@ -565,6 +643,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
data["messages"] = set_last_user_message(messages, sanitized_content)
|
||||
|
||||
except ModelArmorAPIError as e:
|
||||
self._raise_if_fail_closed(e)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -625,7 +705,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
metadata = data.setdefault("metadata", {})
|
||||
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
|
||||
metadata["_model_armor_response"] = self._append_armor_response(
|
||||
metadata.get("_model_armor_response"), armor_response
|
||||
metadata.get("_model_armor_response"),
|
||||
self._build_logging_response(armor_response),
|
||||
)
|
||||
if blocked or metadata.get("_model_armor_status") == "blocked":
|
||||
metadata["_model_armor_status"] = "blocked"
|
||||
|
|
@ -640,10 +721,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if blocked:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Content blocked by Model Armor",
|
||||
"model_armor_response": armor_response,
|
||||
},
|
||||
detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response),
|
||||
)
|
||||
|
||||
# If mask_request_content is enabled, update messages with sanitized content
|
||||
|
|
@ -656,6 +734,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
data["messages"] = set_last_user_message(messages, sanitized_content)
|
||||
|
||||
except ModelArmorAPIError as e:
|
||||
self._raise_if_fail_closed(e)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -698,7 +778,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
# Attach Model Armor response & status to this request's metadata to prevent race conditions
|
||||
if isinstance(armor_response, dict):
|
||||
model_armor_logged_object = {
|
||||
"model_armor_response": armor_response,
|
||||
"model_armor_response": self._build_logging_response(armor_response),
|
||||
"model_armor_status": (
|
||||
"blocked"
|
||||
if self._should_block_content(
|
||||
|
|
@ -729,10 +809,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Response blocked by Model Armor",
|
||||
"model_armor_response": armor_response,
|
||||
},
|
||||
detail=self._build_block_error_detail("Response blocked by Model Armor", armor_response),
|
||||
)
|
||||
|
||||
# If mask_response_content is enabled, update response with sanitized content
|
||||
|
|
@ -746,6 +823,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if choice.message.content:
|
||||
choice.message.content = sanitized_content
|
||||
|
||||
except ModelArmorAPIError as e:
|
||||
self._raise_if_fail_closed(e)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -790,7 +869,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
# Attach Model Armor response & status to this request's metadata to avoid race conditions
|
||||
if isinstance(request_data, dict):
|
||||
metadata = request_data.setdefault("metadata", {})
|
||||
metadata["_model_armor_response"] = armor_response
|
||||
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
|
||||
metadata["_model_armor_status"] = (
|
||||
"blocked" if self._should_block_content(armor_response) else "success"
|
||||
)
|
||||
|
|
@ -809,10 +888,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if self._should_block_content(armor_response):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Streaming response blocked by Model Armor",
|
||||
"model_armor_response": armor_response,
|
||||
},
|
||||
detail=self._build_block_error_detail(
|
||||
"Streaming response blocked by Model Armor",
|
||||
armor_response,
|
||||
),
|
||||
)
|
||||
|
||||
# Apply sanitization if enabled
|
||||
|
|
@ -831,6 +910,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
yield chunk
|
||||
return
|
||||
|
||||
except ModelArmorAPIError as e:
|
||||
if self.optional_params.get("fail_on_error", True):
|
||||
error_obj = {"message": e.detail, "code": "500"}
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
return
|
||||
except HTTPException as e:
|
||||
# Yield error as SSE event so create_response() detects it and
|
||||
# returns a proper JSON error response with the correct status code.
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
aws_sts_endpoint=litellm_params.aws_sts_endpoint,
|
||||
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
|
||||
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
|
||||
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback)
|
||||
return _bedrock_callback
|
||||
|
|
|
|||
|
|
@ -19,7 +19,10 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
|||
iter_client_callback_metadata_dicts,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.litellm_core_utils.url_utils import (
|
||||
is_url_destination_allowed_by_host,
|
||||
provider_url_destination_candidates,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
AddTeamCallback,
|
||||
CommonProxyErrors,
|
||||
|
|
@ -227,23 +230,26 @@ def _reject_url_valued_destinations(data: Dict[str, Any]) -> None:
|
|||
allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
|
||||
for field in _URL_DESTINATION_REQUEST_FIELDS:
|
||||
value = data.get(field)
|
||||
if not isinstance(value, str) or not value.startswith(("http://", "https://")):
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
if is_url_destination_allowed_by_host(value, allowed_hosts):
|
||||
continue
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "invalid_request",
|
||||
"param": field,
|
||||
"message": (
|
||||
f"URL-valued '{field}' is not allowed. Configure custom "
|
||||
"endpoints with api_base instead, or add the destination "
|
||||
"host to `provider_url_destination_allowed_hosts` in "
|
||||
"litellm_settings."
|
||||
),
|
||||
},
|
||||
)
|
||||
for candidate in provider_url_destination_candidates(value):
|
||||
if not candidate.lower().startswith(("http://", "https://")):
|
||||
continue
|
||||
if is_url_destination_allowed_by_host(candidate, allowed_hosts):
|
||||
continue
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "invalid_request",
|
||||
"param": field,
|
||||
"message": (
|
||||
f"URL-valued '{field}' is not allowed. Configure custom "
|
||||
"endpoints with api_base instead, or add the destination "
|
||||
"host to `provider_url_destination_allowed_hosts` in "
|
||||
"litellm_settings."
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _strip_untrusted_request_header_controls(
|
||||
|
|
|
|||
|
|
@ -18,8 +18,9 @@ from pydantic import BaseModel, Field
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._redis import _redis_kwargs_from_environment
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.proxy._types import (
|
||||
AUDIT_ACTIONS,
|
||||
LiteLLM_AuditLogs,
|
||||
|
|
@ -43,6 +44,17 @@ router = APIRouter()
|
|||
# (e.g. redis://:secret@host:6379/1).
|
||||
_CACHE_SENSITIVE_FIELDS: set = {"password", "sentinel_password", "url"}
|
||||
|
||||
# The env fallback resolves the full set of redis.Redis kwargs, which includes
|
||||
# credential-bearing params (azure_client_secret, ssl_password, ...) that are
|
||||
# not cache UI fields. Only overlay fields the settings page actually renders,
|
||||
# so the read never surfaces a credential the UI does not manage.
|
||||
_CACHE_SETTINGS_FIELD_NAMES: frozenset = frozenset(field.field_name for field in CACHE_SETTINGS_FIELDS)
|
||||
|
||||
# Classifier used, alongside _CACHE_SENSITIVE_FIELDS, to redact any
|
||||
# credential-bearing key before it leaves the server (`url` is kept in the
|
||||
# explicit set because its name carries no sensitive segment).
|
||||
_CREDENTIAL_CLASSIFIER = SensitiveDataMasker()
|
||||
|
||||
|
||||
_REDACTED_VALUE = "***REDACTED***"
|
||||
|
||||
|
|
@ -67,6 +79,165 @@ def _resolve_cache_url_precedence(settings: Mapping[str, Any]) -> dict[str, Any]
|
|||
return {k: v for k, v in settings.items() if k not in _URL_OVERRIDDEN_CONNECTION_FIELDS}
|
||||
|
||||
|
||||
def _parse_stored_settings(cache_settings_value: object) -> dict[str, Any]:
|
||||
"""Normalize a stored cache_settings blob to a dict.
|
||||
|
||||
The prisma column comes back as either a JSON string or an already-parsed
|
||||
dict depending on the client, so callers that json.loads unconditionally
|
||||
silently drop the whole (still-encrypted) row on the dict path.
|
||||
"""
|
||||
parsed = json.loads(cache_settings_value) if isinstance(cache_settings_value, str) else cache_settings_value
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def _overlay_environment(stored: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Fill connection fields from the REDIS_* environment the cache actually reads.
|
||||
|
||||
A response cache pointed at Redis resolves host/port/password/etc. from the
|
||||
REDIS_* env vars when the stored config leaves them unset, so a cache
|
||||
configured purely through the environment works while its settings page,
|
||||
which reads only the database row, shows blank. Overlaying the same env
|
||||
kwargs the runtime uses makes the page reflect the effective connection.
|
||||
Stored values win; the environment only fills what the stored config omits.
|
||||
"""
|
||||
env_kwargs = {
|
||||
key: value for key, value in _redis_kwargs_from_environment().items() if key in _CACHE_SETTINGS_FIELD_NAMES
|
||||
}
|
||||
if not env_kwargs:
|
||||
return dict(stored)
|
||||
effective = {**env_kwargs, **stored}
|
||||
# the env fallback is a Redis connection, so name the type when the stored
|
||||
# config did not, letting the UI render the Redis fields it just populated
|
||||
effective.setdefault("type", "redis")
|
||||
return effective
|
||||
|
||||
|
||||
def _redact_credentials(settings: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Replace credential-bearing values with a fixed marker, keeping the rest.
|
||||
|
||||
The marker is unambiguous on the way back in: an admin who edits an
|
||||
unrelated field and re-submits sends the marker for the untouched secret,
|
||||
which the update path maps back to the stored value rather than persisting
|
||||
the marker over a working password.
|
||||
"""
|
||||
return {
|
||||
key: (_REDACTED_VALUE if value is not None and _is_credential_field(key) else value)
|
||||
for key, value in settings.items()
|
||||
}
|
||||
|
||||
|
||||
def _is_credential_field(key: str) -> bool:
|
||||
"""Whether a cache setting carries a credential and must be redacted on read."""
|
||||
return key in _CACHE_SENSITIVE_FIELDS or _CREDENTIAL_CLASSIFIER.is_sensitive_key(key)
|
||||
|
||||
|
||||
def _has_connection_target(value: object) -> bool:
|
||||
"""Whether a payload value names a live discrete connection target."""
|
||||
if isinstance(value, str):
|
||||
return value.strip() != "" and value != _REDACTED_VALUE
|
||||
return value not in (None, [], {})
|
||||
|
||||
|
||||
# Every field that identifies which Redis a credential belongs to, across node
|
||||
# (host/port/url), cluster (redis_startup_nodes), and sentinel
|
||||
# (sentinel_nodes/service_name) modes. A stored secret is bound to these.
|
||||
_CONNECTION_TARGET_FIELDS: tuple = (
|
||||
"host",
|
||||
"port",
|
||||
"url",
|
||||
"redis_startup_nodes",
|
||||
"sentinel_nodes",
|
||||
"service_name",
|
||||
)
|
||||
|
||||
|
||||
def _target_repr(value: object) -> str:
|
||||
"""Canonical string form of a connection-target value for equality checks.
|
||||
|
||||
The client may serialize the same target differently from storage (a port as
|
||||
"6379" vs 6379, node lists round-tripped through JSON), so compare normalized
|
||||
forms rather than raw values to avoid treating an unchanged target as a change.
|
||||
"""
|
||||
if isinstance(value, (list, dict)):
|
||||
return json.dumps(value, sort_keys=True, default=str)
|
||||
return str(value)
|
||||
|
||||
|
||||
def _saved_secret_is_reusable(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> bool:
|
||||
"""Whether a stored credential may be restored for this request.
|
||||
|
||||
A stored secret belongs to the stored connection target, so it is reused only
|
||||
when the request describes that same target on every dimension the stored
|
||||
config pins (host/port, url, cluster nodes, sentinel nodes/service). This
|
||||
prevents credential replay: a caller cannot omit the credential, point at a
|
||||
different (or incomplete) target, and have the proxy send the stored secret
|
||||
to a Redis of their choosing.
|
||||
|
||||
Non-secret target fields (host/port/nodes/service) must be supplied and match
|
||||
in normalized form, so equivalent representations (port "6379" vs 6379) are
|
||||
not seen as a change while an omitted or different value is. ``url`` is the
|
||||
exception: it is itself the secret and the form never re-prefills it, so a
|
||||
redacted or omitted url means "keep the stored url" (same target) and only a
|
||||
different supplied url blocks reuse.
|
||||
"""
|
||||
for field in _CONNECTION_TARGET_FIELDS:
|
||||
saved_value = saved.get(field)
|
||||
if saved_value in (None, "", [], {}):
|
||||
continue # the stored config does not pin this dimension
|
||||
incoming_value = incoming.get(field)
|
||||
if field == "url":
|
||||
if incoming_value in (None, "", _REDACTED_VALUE):
|
||||
continue # url kept as-is (same target)
|
||||
if _target_repr(incoming_value) != _target_repr(saved_value):
|
||||
return False
|
||||
continue
|
||||
if _target_repr(incoming_value) != _target_repr(saved_value):
|
||||
return False # a pinned target field is missing or different
|
||||
return True
|
||||
|
||||
|
||||
def _merge_over_saved(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Keep the stored secret behind any credential the caller echoed back redacted or omitted.
|
||||
|
||||
GET returns credentials as the marker and the form never re-prefills a
|
||||
secret, so a save that does not touch a credential arrives with the marker
|
||||
or with the field absent. Either way the real secret must survive: it is
|
||||
restored from the stored row, or dropped when there is no stored row (the
|
||||
value is env-sourced and the marker must never be persisted). Non-secret
|
||||
fields are taken from the incoming payload as-is, so clearing one still works.
|
||||
|
||||
``url`` is the exception: it is credential-bearing (redacted) yet also a
|
||||
connection-mode selector that url-precedence resolves against host/port. If
|
||||
the caller supplies a discrete target (host, cluster, or sentinel nodes), a
|
||||
stored url is a stale mode the caller is leaving, so it is dropped rather
|
||||
than restored, otherwise url-precedence would resurrect it and discard the
|
||||
submitted host/port.
|
||||
"""
|
||||
switching_to_discrete_target = (
|
||||
_has_connection_target(incoming.get("host"))
|
||||
or _has_connection_target(incoming.get("redis_startup_nodes"))
|
||||
or _has_connection_target(incoming.get("sentinel_nodes"))
|
||||
)
|
||||
reuse_saved_secret = _saved_secret_is_reusable(incoming, saved)
|
||||
merged = dict(incoming)
|
||||
for field in _CACHE_SENSITIVE_FIELDS:
|
||||
# A value the caller explicitly supplied is honored verbatim: a new
|
||||
# secret, or an empty string / null to clear the stored one. Only an
|
||||
# omitted field or the echoed-back marker triggers preserve-or-drop.
|
||||
if field in incoming and incoming[field] != _REDACTED_VALUE:
|
||||
continue
|
||||
if field == "url" and switching_to_discrete_target:
|
||||
merged.pop(field, None)
|
||||
continue
|
||||
if field in saved and reuse_saved_secret:
|
||||
merged[field] = saved[field]
|
||||
else:
|
||||
# nothing stored to reuse, or the caller is pointing at a different
|
||||
# target: never persist/replay the marker or the stored secret
|
||||
merged.pop(field, None)
|
||||
return merged
|
||||
|
||||
|
||||
def _redact_settings(settings: Optional[Mapping[str, Any]]) -> Dict[str, Any]:
|
||||
"""Replace every value in a settings map with a fixed marker.
|
||||
|
||||
|
|
@ -270,34 +441,34 @@ async def get_cache_settings(
|
|||
# Get cache settings fields from types file
|
||||
cache_fields = [field.model_copy(deep=True) for field in CACHE_SETTINGS_FIELDS]
|
||||
|
||||
# Try to get cache settings from database
|
||||
current_values = {}
|
||||
# Read the stored settings (decrypted); an env-only cache has none.
|
||||
stored: dict[str, Any] = {}
|
||||
if prisma_client is not None:
|
||||
cache_config = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
# Decrypt cache settings
|
||||
cache_settings_json = cache_config.cache_settings
|
||||
if isinstance(cache_settings_json, str):
|
||||
cache_settings_dict = json.loads(cache_settings_json)
|
||||
else:
|
||||
cache_settings_dict = cache_settings_json
|
||||
stored = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(cache_config.cache_settings)
|
||||
)
|
||||
|
||||
# Decrypt environment variables
|
||||
decrypted_settings = proxy_config._decrypt_db_variables(variables_dict=cache_settings_dict)
|
||||
# Fill connection fields from the REDIS_* environment the cache resolves
|
||||
# from when the stored config leaves them unset, then apply url precedence
|
||||
# so a url-mode config does not surface conflicting discrete fields (which
|
||||
# would otherwise let a no-op save silently switch it to host/port).
|
||||
effective = _resolve_cache_url_precedence(_overlay_environment(stored))
|
||||
|
||||
# Derive redis_type for UI based on settings
|
||||
# UI uses redis_type to show/hide fields, backend only stores 'type'
|
||||
if decrypted_settings.get("type") == "redis":
|
||||
if decrypted_settings.get("redis_startup_nodes"):
|
||||
decrypted_settings["redis_type"] = "cluster"
|
||||
elif decrypted_settings.get("sentinel_nodes"):
|
||||
decrypted_settings["redis_type"] = "sentinel"
|
||||
else:
|
||||
decrypted_settings["redis_type"] = "node"
|
||||
# Derive redis_type for UI based on settings
|
||||
# UI uses redis_type to show/hide fields, backend only stores 'type'
|
||||
if effective.get("type") == "redis":
|
||||
if effective.get("redis_startup_nodes"):
|
||||
effective["redis_type"] = "cluster"
|
||||
elif effective.get("sentinel_nodes"):
|
||||
effective["redis_type"] = "sentinel"
|
||||
else:
|
||||
effective["redis_type"] = "node"
|
||||
|
||||
# Mask credential fields so the GET response never carries
|
||||
# plaintext Redis / Sentinel passwords off the server.
|
||||
current_values = mask_sensitive_keys(decrypted_settings, _CACHE_SENSITIVE_FIELDS)
|
||||
# Redact credential fields so the GET response never carries a plaintext
|
||||
# Redis / Sentinel password off the server.
|
||||
current_values = _redact_credentials(effective)
|
||||
|
||||
# Update field values with current values
|
||||
for field in cache_fields:
|
||||
|
|
@ -331,10 +502,27 @@ async def test_cache_connection(
|
|||
to verify the credentials work without affecting global state.
|
||||
"""
|
||||
from litellm import Cache
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_config
|
||||
|
||||
try:
|
||||
cache_settings = _resolve_cache_url_precedence(request.cache_settings)
|
||||
verbose_proxy_logger.debug("Testing cache connection with settings: %s", cache_settings)
|
||||
# A credential the form left untouched arrives redacted; resolve it back
|
||||
# to the stored secret so the test connects with the real password. A
|
||||
# lookup failure must not block the test, so fall back to no stored row.
|
||||
saved_settings: dict[str, Any] = {}
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
existing_row = await CacheConfigRepository(prisma_client).table.find_unique(
|
||||
where={"id": "cache_config"}
|
||||
)
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
saved_settings = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(existing_row.cache_settings)
|
||||
)
|
||||
except Exception: # noqa: BLE001 - a saved-settings lookup failure must not block a connection test
|
||||
saved_settings = {}
|
||||
cache_settings = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings))
|
||||
# cache_settings now carries the resolved plaintext credential; never log it raw
|
||||
verbose_proxy_logger.debug("Testing cache connection with settings: %s", _redact_credentials(cache_settings))
|
||||
|
||||
# Only support Redis for now
|
||||
if cache_settings.get("type") != "redis":
|
||||
|
|
@ -400,19 +588,20 @@ async def update_cache_settings(
|
|||
)
|
||||
|
||||
try:
|
||||
cache_settings = _resolve_cache_url_precedence(request.cache_settings)
|
||||
|
||||
# Snapshot the prior settings (key set only — values get redacted in
|
||||
# the audit row) so the audit-log entry shows which fields changed.
|
||||
# Read the stored row first: its decrypted values back any credential the
|
||||
# caller echoed back redacted, and its key set drives the audit diff.
|
||||
existing_row = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
before_settings: Optional[Dict[str, Any]] = None
|
||||
saved_settings: dict[str, Any] = {}
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
try:
|
||||
before_settings = json.loads(existing_row.cache_settings)
|
||||
except (TypeError, ValueError):
|
||||
before_settings = None
|
||||
before_settings = _parse_stored_settings(existing_row.cache_settings)
|
||||
saved_settings = proxy_config._decrypt_db_variables(variables_dict=before_settings)
|
||||
action: AUDIT_ACTIONS = "updated" if existing_row is not None else "created"
|
||||
|
||||
# Preserve stored secrets behind any redacted or omitted credential, then
|
||||
# resolve the url-vs-discrete-fields precedence.
|
||||
cache_settings = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings))
|
||||
|
||||
# Encrypt sensitive fields (keep redis_type for storage)
|
||||
encrypted_settings = proxy_config._encrypt_env_variables(environment_variables=cache_settings)
|
||||
|
||||
|
|
@ -461,7 +650,7 @@ async def update_cache_settings(
|
|||
return {
|
||||
"message": "Cache settings updated successfully",
|
||||
"status": "success",
|
||||
"settings": cache_settings,
|
||||
"settings": _redact_credentials(cache_settings),
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error updating cache settings: {str(e)}")
|
||||
|
|
|
|||
|
|
@ -42,6 +42,9 @@ from litellm.proxy._experimental.mcp_server.db import (
|
|||
rotate_mcp_user_credentials_master_key,
|
||||
rotate_mcp_user_env_vars_master_key,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
rotate_sso_identity_assertions_master_key,
|
||||
)
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, hash_token
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
|
|
@ -4242,6 +4245,15 @@ async def _rotate_master_key(
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to rotate MCP user env vars: %s", str(e))
|
||||
|
||||
# 4d. process SSO identity assertion table (EMA subject tokens)
|
||||
try:
|
||||
await rotate_sso_identity_assertions_master_key(
|
||||
prisma_client=prisma_client,
|
||||
new_master_key=new_master_key,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation
|
||||
verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e))
|
||||
|
||||
# 5. process credentials table
|
||||
try:
|
||||
credentials = await CredentialsRepository(prisma_client).table.find_many()
|
||||
|
|
|
|||
|
|
@ -136,9 +136,12 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_raise_if_not_oauth2,
|
||||
authorize_with_server,
|
||||
client_supplied_redirect_uris,
|
||||
exchange_token_with_server,
|
||||
get_request_base_url,
|
||||
redeem_passthrough_authorization_code,
|
||||
register_client_with_server,
|
||||
resolve_ephemeral_dcr_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -1661,7 +1664,21 @@ if MCP_AVAILABLE:
|
|||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
# Use the server's stored client_id when the caller doesn't supply one
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
stored_or_supplied_client_id = mcp_server.client_id or client_id or ""
|
||||
ephemeral_dcr_client = (
|
||||
await resolve_ephemeral_dcr_client(
|
||||
request=request,
|
||||
mcp_server=mcp_server,
|
||||
code_challenge=code_challenge,
|
||||
code_challenge_method=code_challenge_method,
|
||||
redirect_uri=redirect_uri,
|
||||
)
|
||||
if not stored_or_supplied_client_id
|
||||
else None
|
||||
)
|
||||
resolved_client_id = stored_or_supplied_client_id or (
|
||||
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
|
||||
)
|
||||
if not resolved_client_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
|
|
@ -1683,6 +1700,7 @@ if MCP_AVAILABLE:
|
|||
code_challenge_method=code_challenge_method,
|
||||
response_type=response_type,
|
||||
scope=scope,
|
||||
ephemeral_dcr_client=ephemeral_dcr_client,
|
||||
)
|
||||
|
||||
@router.post(
|
||||
|
|
@ -1705,7 +1723,21 @@ if MCP_AVAILABLE:
|
|||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
|
||||
# grant must never open one: the minted client is unrecoverable after the single flow by
|
||||
# contract, so an expired browser-held token re-runs authorize instead.
|
||||
sealed_code = (
|
||||
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
|
||||
if grant_type == "authorization_code"
|
||||
else None
|
||||
)
|
||||
resolved_code = sealed_code.upstream_code if sealed_code else code
|
||||
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
|
||||
# or plain flow alike), so the exchange must present that binding, not the browser page.
|
||||
resolved_redirect_uri = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
|
||||
caller_client_id = sealed_code.client_id if sealed_code else client_id
|
||||
caller_client_secret = sealed_code.client_secret if sealed_code else client_secret
|
||||
resolved_client_id = mcp_server.client_id or caller_client_id or ""
|
||||
if not resolved_client_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
|
|
@ -1721,13 +1753,14 @@ if MCP_AVAILABLE:
|
|||
request=request,
|
||||
mcp_server=mcp_server,
|
||||
grant_type=grant_type,
|
||||
code=code,
|
||||
redirect_uri=redirect_uri,
|
||||
code=resolved_code,
|
||||
redirect_uri=resolved_redirect_uri,
|
||||
client_id=resolved_client_id,
|
||||
client_secret=client_secret,
|
||||
client_secret=caller_client_secret,
|
||||
code_verifier=code_verifier,
|
||||
refresh_token=refresh_token,
|
||||
scope=scope,
|
||||
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
|
||||
)
|
||||
|
||||
@router.post(
|
||||
|
|
@ -1743,6 +1776,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
request_data = await _read_request_body(request=request)
|
||||
data: dict = {**request_data}
|
||||
client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris"))
|
||||
|
||||
return await register_client_with_server(
|
||||
request=request,
|
||||
|
|
@ -1753,6 +1787,7 @@ if MCP_AVAILABLE:
|
|||
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
|
||||
fallback_client_id=server_id,
|
||||
persist_credentials=_user_is_full_admin(user_api_key_dict),
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
)
|
||||
|
||||
@router.delete(
|
||||
|
|
|
|||
|
|
@ -13,16 +13,18 @@ Endpoints for /organization operations
|
|||
|
||||
#### ORGANIZATION MANAGEMENT ####
|
||||
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import can_user_call_model, get_user_object
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.management_endpoints.budget_management_endpoints import (
|
||||
new_budget,
|
||||
update_budget,
|
||||
|
|
@ -34,6 +36,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
handle_update_object_permission_common,
|
||||
prepare_object_permission_upsert,
|
||||
)
|
||||
from litellm.proxy.management_helpers.utils import (
|
||||
get_new_internal_user_defaults,
|
||||
|
|
@ -101,6 +104,30 @@ async def _verify_org_access(
|
|||
)
|
||||
|
||||
|
||||
_STR_OBJECT_DICT_ADAPTER = TypeAdapter(dict[str, object])
|
||||
_BUDGET_SETTABLE_FIELDS = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"}
|
||||
_ORG_COLUMN_FIELDS = frozenset({"organization_alias", "models"})
|
||||
|
||||
|
||||
def build_budget_write_data(budget_updates: Mapping[str, object], updated_by: str) -> Mapping[str, object]:
|
||||
"""
|
||||
Budget-row columns to write. ``budget_reset_at`` tracks any sent ``budget_duration``:
|
||||
recomputed for a new duration, cleared alongside a ``None`` duration so no stale reset
|
||||
timestamp survives. Other sent fields (including a ``None`` clear) are written as-is.
|
||||
"""
|
||||
budget_duration = budget_updates.get("budget_duration")
|
||||
recomputed_reset_at: Mapping[str, object] = (
|
||||
{
|
||||
"budget_reset_at": (
|
||||
get_budget_reset_time(budget_duration=budget_duration) if isinstance(budget_duration, str) else None
|
||||
)
|
||||
}
|
||||
if "budget_duration" in budget_updates
|
||||
else {}
|
||||
)
|
||||
return {**budget_updates, **recomputed_reset_at, "updated_by": updated_by}
|
||||
|
||||
|
||||
def handle_nested_budget_structure_in_organization_update_request(
|
||||
raw_data: dict,
|
||||
) -> dict:
|
||||
|
|
@ -556,6 +583,154 @@ async def handle_update_object_permission(
|
|||
return data_json
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/v2/organization/{organization_id}",
|
||||
tags=["organization management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=LiteLLM_OrganizationTableWithMembers,
|
||||
include_in_schema=False,
|
||||
)
|
||||
async def update_organization_v2(
|
||||
organization_id: str,
|
||||
data: OrganizationUpdateRequestV2,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
):
|
||||
"""
|
||||
Partial update of an organization (RESTful PATCH, RFC 7396 merge-patch semantics).
|
||||
|
||||
A sent field is written and an omitted one is left untouched (presence is read from
|
||||
``model_fields_set``). Clear tokens are per field: budget limits and ``metadata`` clear with
|
||||
``null``, ``models`` with ``[]``, and ``object_permission`` with ``null`` (it merges when sent,
|
||||
so an empty ``{}`` is rejected). ``organization_alias`` is required and cannot be cleared.
|
||||
Validation failures return 422; the object-permission upsert, budget-row write, and
|
||||
org-row write are one transaction.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
if user_api_key_dict.user_id is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Cannot associate a user_id to this action. Check `/key/info` to validate if 'user_id' is set."
|
||||
},
|
||||
)
|
||||
|
||||
if data.max_budget is not None and (not math.isfinite(data.max_budget) or data.max_budget < 0):
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"},
|
||||
)
|
||||
if data.soft_budget is not None and (not math.isfinite(data.soft_budget) or data.soft_budget < 0):
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
|
||||
)
|
||||
if data.model_max_budget:
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
validate_model_max_budget,
|
||||
)
|
||||
|
||||
try:
|
||||
validate_model_max_budget(data.model_max_budget)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=422, detail={"error": str(e)})
|
||||
|
||||
if "organization_alias" in data.model_fields_set and data.organization_alias is None:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail={"error": "organization_alias cannot be cleared; it is required"},
|
||||
)
|
||||
if "models" in data.model_fields_set and data.models is None:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail={"error": "models cannot be set to null; send [] to clear it"},
|
||||
)
|
||||
if data.object_permission is not None and not data.object_permission.model_dump(exclude_none=True):
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail={
|
||||
"error": "object_permission cannot be an empty object; send null to clear it, or a non-empty object to set grants"
|
||||
},
|
||||
)
|
||||
|
||||
await _verify_org_access(
|
||||
organization_id=organization_id,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
where={"organization_id": organization_id},
|
||||
)
|
||||
if existing_organization_row is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Organization not found for organization_id={organization_id}"},
|
||||
)
|
||||
|
||||
field_values = _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump())
|
||||
present_fields = data.model_fields_set
|
||||
budget_updates = {field: field_values[field] for field in present_fields if field in _BUDGET_SETTABLE_FIELDS}
|
||||
org_column_updates: Mapping[str, object] = {
|
||||
**{field: field_values[field] for field in present_fields if field in _ORG_COLUMN_FIELDS},
|
||||
**({"metadata": data.metadata or {}} if "metadata" in present_fields else {}),
|
||||
}
|
||||
|
||||
object_permission_cleared = "object_permission" in present_fields and data.object_permission is None
|
||||
object_permission_upsert = (
|
||||
await prepare_object_permission_upsert(
|
||||
new_object_permission=data.object_permission.model_dump(exclude_none=True),
|
||||
existing_object_permission_id=existing_organization_row.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if data.object_permission is not None
|
||||
else None
|
||||
)
|
||||
object_permission_write: Mapping[str, object] = (
|
||||
{"object_permission_id": object_permission_upsert.object_permission_id}
|
||||
if object_permission_upsert is not None
|
||||
else ({"object_permission_id": None} if object_permission_cleared else {})
|
||||
)
|
||||
|
||||
organization_write_data = prisma_client.jsonify_object(
|
||||
{
|
||||
**org_column_updates,
|
||||
**object_permission_write,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
)
|
||||
|
||||
async with prisma_client.db.tx() as tx:
|
||||
if object_permission_upsert is not None:
|
||||
await tx.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": object_permission_upsert.object_permission_id},
|
||||
data={
|
||||
"create": object_permission_upsert.record,
|
||||
"update": object_permission_upsert.record,
|
||||
},
|
||||
)
|
||||
if budget_updates:
|
||||
await tx.litellm_budgettable.update(
|
||||
where={"budget_id": existing_organization_row.budget_id},
|
||||
data=prisma_client.jsonify_object(
|
||||
dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id))
|
||||
),
|
||||
)
|
||||
response = await tx.litellm_organizationtable.update(
|
||||
where={"organization_id": organization_id},
|
||||
data=organization_write_data,
|
||||
include={"members": True, "teams": True, "litellm_budget_table": True},
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/organization/delete",
|
||||
tags=["organization management"],
|
||||
|
|
|
|||
|
|
@ -98,14 +98,27 @@ class UserProvisionerHelpers:
|
|||
if not existing_user:
|
||||
return None
|
||||
|
||||
# Update the user
|
||||
new_teams = list(dict.fromkeys(new_user_request.teams or []))
|
||||
|
||||
if new_user_request.user_id != existing_user.user_id:
|
||||
await UserRepository(prisma_client).table.update(
|
||||
where={"user_id": existing_user.user_id},
|
||||
data={"user_id": new_user_request.user_id},
|
||||
)
|
||||
|
||||
await _handle_team_membership_changes(
|
||||
user_id=new_user_request.user_id,
|
||||
existing_teams=existing_user.teams or [],
|
||||
new_teams=new_teams,
|
||||
raise_on_error=True,
|
||||
)
|
||||
|
||||
updated_user = await UserRepository(prisma_client).table.update(
|
||||
where={"user_id": existing_user.user_id},
|
||||
where={"user_id": new_user_request.user_id},
|
||||
data={
|
||||
"user_id": new_user_request.user_id,
|
||||
"user_email": new_user_request.user_email,
|
||||
"user_alias": new_user_request.user_alias,
|
||||
"teams": new_user_request.teams,
|
||||
"teams": new_teams,
|
||||
"metadata": safe_dumps(new_user_request.metadata),
|
||||
**({"user_role": new_user_request.user_role} if admin_group is not None else {}),
|
||||
},
|
||||
|
|
@ -440,7 +453,12 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]:
|
|||
return members
|
||||
|
||||
|
||||
async def _handle_team_membership_changes(user_id: str, existing_teams: List[str], new_teams: List[str]) -> None:
|
||||
async def _handle_team_membership_changes(
|
||||
user_id: str,
|
||||
existing_teams: List[str],
|
||||
new_teams: List[str],
|
||||
raise_on_error: bool = False,
|
||||
) -> None:
|
||||
"""Handle adding/removing user from teams based on changes."""
|
||||
existing_teams_set = set(existing_teams)
|
||||
new_teams_set = set(new_teams)
|
||||
|
|
@ -453,6 +471,7 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str
|
|||
user_id=user_id,
|
||||
teams_ids_to_add_user_to=list(teams_to_add),
|
||||
teams_ids_to_remove_user_from=list(teams_to_remove),
|
||||
raise_on_error=raise_on_error,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1298,6 +1317,13 @@ async def delete_user(
|
|||
where={"team_id": team.team_id}, data={"members": new_members}
|
||||
)
|
||||
|
||||
team_row = LiteLLM_TeamTable(**team.model_dump())
|
||||
if any(member.user_id == user_id for member in team_row.members_with_roles or []):
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
await _set_user_keys_blocked(user_id=user_id, blocked=True)
|
||||
|
||||
await _delete_rows_referencing_user(prisma_client, user_id=user_id)
|
||||
|
|
@ -1327,6 +1353,31 @@ def _extract_group_values(value: Any) -> List[str]:
|
|||
return group_values
|
||||
|
||||
|
||||
def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]:
|
||||
"""Return ids from a SCIM filtered path like ``members[value eq "id"]``.
|
||||
|
||||
Okta commonly sends membership removals as a filtered path and omits the
|
||||
request body ``value``, so the id lives only inside the ``[value eq "..."]``
|
||||
filter. The ``eq`` operator is matched case-insensitively per the SCIM
|
||||
spec; the id keeps its original case. Per the SCIM filter grammar the
|
||||
compared value must be quoted (single or double), so malformed unquoted
|
||||
filters yield no id. A quoted id may contain escaped quotes and
|
||||
backslashes (``\\"`` and ``\\\\``), which are unescaped before use.
|
||||
``path`` must be the raw, case-preserving path from the patch op.
|
||||
"""
|
||||
if not path:
|
||||
return []
|
||||
match = re.match(
|
||||
rf"""\s*{re.escape(attribute)}\s*\[\s*value\s+eq\s+(['"])((?:\\.|[^\\])*?)\1\s*\]\s*$""",
|
||||
path,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
if not match:
|
||||
return []
|
||||
extracted = re.sub(r"\\(.)", r"\1", match.group(2))
|
||||
return [extracted] if extracted else []
|
||||
|
||||
|
||||
def _handle_displayname_update(op_type: str, value: Any, update_data: Dict[str, Any]) -> None:
|
||||
"""Handle displayname updates."""
|
||||
if op_type == "remove":
|
||||
|
|
@ -1370,9 +1421,11 @@ def _handle_name_update(path: str, op_type: str, value: Any, scim_metadata: Dict
|
|||
scim_metadata["familyName"] = str(value)
|
||||
|
||||
|
||||
def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str]) -> Optional[Set[str]]:
|
||||
def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str], path: str | None) -> Set[str] | None:
|
||||
"""Handle group/team membership operations."""
|
||||
group_values = _extract_group_values(value)
|
||||
if not group_values and value is None:
|
||||
group_values = _extract_ids_from_path_filter(path, "groups")
|
||||
if op_type == "replace":
|
||||
return set(group_values)
|
||||
elif op_type == "add":
|
||||
|
|
@ -1485,7 +1538,7 @@ def _apply_patch_ops(
|
|||
elif _multi_valued_attribute_base(path) in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS:
|
||||
_handle_multi_valued_attribute_update(path, op_type, value, metadata)
|
||||
elif path.startswith("groups"):
|
||||
new_replace_set = _handle_group_operations(op_type, value, teams_set)
|
||||
new_replace_set = _handle_group_operations(op_type, value, teams_set, op.path)
|
||||
if new_replace_set is not None:
|
||||
replace_team_set = new_replace_set
|
||||
else:
|
||||
|
|
@ -1497,16 +1550,29 @@ def _apply_patch_ops(
|
|||
return update_data, final_team_set
|
||||
|
||||
|
||||
def _is_user_not_in_team_error(exc: HTTPException) -> bool:
|
||||
"""True when team_member_delete reports the user was already absent from the
|
||||
team, which is the idempotent no-op case for a removal."""
|
||||
detail = exc.detail
|
||||
return isinstance(detail, dict) and detail.get("error") == "User not found in team"
|
||||
|
||||
|
||||
async def patch_team_membership(
|
||||
user_id: str,
|
||||
teams_ids_to_add_user_to: List[str],
|
||||
teams_ids_to_remove_user_from: List[str],
|
||||
raise_on_error: bool = False,
|
||||
) -> bool:
|
||||
"""
|
||||
Add or remove user from teams
|
||||
|
||||
Handles duplicate membership gracefully (idempotent operation).
|
||||
If a user is already in a team, that's fine - we don't treat it as an error.
|
||||
A user already being in a team (on add) or already absent from it (on
|
||||
remove) is treated as a no-op, not an error.
|
||||
|
||||
When ``raise_on_error`` is True a genuine add or remove failure (anything
|
||||
other than those idempotent no-ops) propagates instead of being swallowed,
|
||||
so a caller can avoid persisting a teams array the roster never received.
|
||||
"""
|
||||
for _team_id in teams_ids_to_add_user_to:
|
||||
try:
|
||||
|
|
@ -1521,9 +1587,13 @@ async def patch_team_membership(
|
|||
# Handle duplicate membership gracefully - this is idempotent
|
||||
if e.type == ProxyErrorTypes.team_member_already_in_team:
|
||||
verbose_proxy_logger.debug(f"User {user_id} is already in team {_team_id}, skipping add")
|
||||
elif raise_on_error:
|
||||
raise
|
||||
else:
|
||||
verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}")
|
||||
except Exception as e:
|
||||
if raise_on_error:
|
||||
raise
|
||||
verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}")
|
||||
|
||||
for _team_id in teams_ids_to_remove_user_from:
|
||||
|
|
@ -1532,7 +1602,16 @@ async def patch_team_membership(
|
|||
data=TeamMemberDeleteRequest(team_id=_team_id, user_id=user_id),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
except HTTPException as e:
|
||||
if _is_user_not_in_team_error(e):
|
||||
verbose_proxy_logger.debug(f"User {user_id} is not in team {_team_id}, skipping remove")
|
||||
elif raise_on_error:
|
||||
raise
|
||||
else:
|
||||
verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}")
|
||||
except Exception as e:
|
||||
if raise_on_error:
|
||||
raise
|
||||
verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}")
|
||||
|
||||
return True
|
||||
|
|
@ -1654,8 +1733,11 @@ async def get_groups(
|
|||
# Convert to SCIM format
|
||||
scim_groups = []
|
||||
for team in teams:
|
||||
# Get team members with display names
|
||||
members = await _get_team_members_display(team.members or [])
|
||||
# Get team members with display names. members_with_roles is the
|
||||
# source of truth; the legacy `members` column is not populated by
|
||||
# team creation, so reading it here would report an empty member
|
||||
# list to the IdP and trigger repeated re-provisioning.
|
||||
members = await _get_team_members_display(await _get_team_member_user_ids_from_team(team))
|
||||
verbose_proxy_logger.debug(f"SCIM GET GROUPS members: {members}")
|
||||
team_alias = getattr(team, "team_alias", team.team_id)
|
||||
team_created_at = team.created_at.isoformat() if team.created_at else None
|
||||
|
|
@ -1877,16 +1959,28 @@ async def delete_group(
|
|||
|
||||
async def _process_group_patch_operations(
|
||||
patch_ops: SCIMPatchOp, existing_team, prisma_client
|
||||
) -> Tuple[Dict[str, Any], Set[str]]:
|
||||
"""Process patch operations for a group and return update data and final members."""
|
||||
) -> Tuple[Dict[str, Any], Set[str], Set[str] | None]:
|
||||
"""Process patch operations for a group and return update data, final members
|
||||
and, when the request contained a member ``replace`` op, the absolute target
|
||||
roster it declared (``None`` otherwise).
|
||||
|
||||
``add``/``remove`` are deltas relative to the current roster, but ``replace``
|
||||
is absolute: it declares the roster is exactly this set, so the caller must
|
||||
reconcile against it as a set-to-target rather than rebasing it onto a
|
||||
concurrently-mutated roster.
|
||||
"""
|
||||
update_data: Dict[str, Any] = {}
|
||||
|
||||
# Create a fresh copy of existing metadata to avoid Prisma issues
|
||||
existing_metadata = existing_team.metadata or {}
|
||||
metadata = dict(existing_metadata) if existing_metadata else {}
|
||||
|
||||
# Track member changes
|
||||
current_members = set(existing_team.members or [])
|
||||
# Track member changes. members_with_roles is the source of truth for team
|
||||
# membership; the legacy `members` column is not populated by team creation
|
||||
# or the real team endpoints, so seeding from it would make an `add`/`remove`
|
||||
# operation recompute the member set from an empty base and silently drop
|
||||
# everyone already in the team.
|
||||
current_members = set(await _get_team_member_user_ids_from_team(existing_team))
|
||||
final_members = current_members.copy()
|
||||
|
||||
# Process each patch operation
|
||||
|
|
@ -1908,6 +2002,8 @@ async def _process_group_patch_operations(
|
|||
elif path.startswith("members"):
|
||||
# Handle member operations
|
||||
member_values = _extract_group_values(value)
|
||||
if not member_values and value is None:
|
||||
member_values = _extract_ids_from_path_filter(op.path, "members")
|
||||
# Check the feature flag
|
||||
scim_upsert_user = await _get_scim_upsert_user_setting()
|
||||
# Validate all users exist or create them based on feature flag
|
||||
|
|
@ -1960,27 +2056,32 @@ async def _process_group_patch_operations(
|
|||
if metadata:
|
||||
update_data["metadata"] = metadata
|
||||
|
||||
return update_data, final_members
|
||||
member_replace_present = any(
|
||||
op.op == "replace" and (op.path or "").lower().startswith("members") for op in patch_ops.Operations
|
||||
)
|
||||
replace_target = set(final_members) if member_replace_present else None
|
||||
|
||||
return update_data, final_members, replace_target
|
||||
|
||||
|
||||
async def _apply_group_patch_updates(
|
||||
group_id: str, update_data: Dict[str, Any], final_members: Set[str], prisma_client
|
||||
):
|
||||
"""Apply patch updates to the group in the database."""
|
||||
# Serialize metadata if present
|
||||
async def _apply_group_patch_updates(group_id: str, update_data: Dict[str, Any], prisma_client):
|
||||
"""Apply the group's metadata/displayName patch updates to the database.
|
||||
|
||||
Membership itself is not written here; it is reconciled onto the source of
|
||||
truth (members_with_roles and each member's user.teams) by
|
||||
_handle_group_membership_changes via team_member_add/team_member_delete.
|
||||
Writing the legacy `members` column here too would create a second, unread
|
||||
copy of membership that could drift from the source of truth.
|
||||
"""
|
||||
if "metadata" in update_data and isinstance(update_data["metadata"], dict):
|
||||
update_data["metadata"] = safe_dumps(update_data["metadata"])
|
||||
|
||||
# Update members list
|
||||
update_data["members"] = list(final_members)
|
||||
|
||||
# Update team in database
|
||||
updated_team = await TeamRepository(prisma_client).table.update(
|
||||
where={"team_id": group_id},
|
||||
data=update_data,
|
||||
)
|
||||
|
||||
return updated_team
|
||||
if update_data:
|
||||
return await TeamRepository(prisma_client).table.update(
|
||||
where={"team_id": group_id},
|
||||
data=update_data,
|
||||
)
|
||||
return await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id})
|
||||
|
||||
|
||||
async def _handle_group_membership_changes(group_id: str, current_members: Set[str], final_members: Set[str]):
|
||||
|
|
@ -2031,27 +2132,29 @@ async def patch_group(
|
|||
existing_team = await _check_team_exists(group_id)
|
||||
|
||||
# Process patch operations
|
||||
update_data, final_members = await _process_group_patch_operations(patch_ops, existing_team, prisma_client)
|
||||
update_data, final_members, replace_target = await _process_group_patch_operations(
|
||||
patch_ops, existing_team, prisma_client
|
||||
)
|
||||
|
||||
# Track current members BEFORE update for comparison
|
||||
current_members = set(await _get_team_member_user_ids_from_team(existing_team))
|
||||
snapshot_members = set(await _get_team_member_user_ids_from_team(existing_team))
|
||||
intended_add = final_members - snapshot_members
|
||||
intended_remove = snapshot_members - final_members
|
||||
|
||||
# Apply updates to the database
|
||||
updated_team = await _apply_group_patch_updates(group_id, update_data, final_members, prisma_client)
|
||||
# Apply the metadata/displayName updates to the database
|
||||
updated_team = await _apply_group_patch_updates(group_id, update_data, prisma_client)
|
||||
|
||||
# Refresh team data from database to get the latest state after concurrent updates
|
||||
# This prevents race conditions when multiple PATCH requests come in simultaneously
|
||||
refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id})
|
||||
if refreshed_team:
|
||||
# Re-read current members from refreshed team to account for concurrent updates
|
||||
refreshed_current_members = set(
|
||||
await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump()))
|
||||
)
|
||||
# Use the refreshed members for comparison
|
||||
current_members = refreshed_current_members
|
||||
refreshed_current = (
|
||||
set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump())))
|
||||
if refreshed_team
|
||||
else snapshot_members
|
||||
)
|
||||
|
||||
# Handle user-team relationship changes
|
||||
await _handle_group_membership_changes(group_id, current_members, final_members)
|
||||
effective_final = (
|
||||
replace_target if replace_target is not None else (refreshed_current | intended_add) - intended_remove
|
||||
)
|
||||
|
||||
await _handle_group_membership_changes(group_id, refreshed_current, effective_final)
|
||||
|
||||
# A rename can flip whether this group matches scim_admin_group by display
|
||||
# name, so retained members must be re-resolved too, not just the ones whose
|
||||
|
|
@ -2060,7 +2163,7 @@ async def patch_group(
|
|||
alias_changed = new_alias != existing_team.team_alias
|
||||
await _recompute_scim_member_roles(
|
||||
prisma_client,
|
||||
(current_members | final_members if alias_changed else current_members ^ final_members),
|
||||
(refreshed_current | effective_final if alias_changed else refreshed_current ^ effective_final),
|
||||
)
|
||||
|
||||
# Refresh team one more time to get final state after membership changes
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ from litellm.proxy._types import (
|
|||
Member,
|
||||
NewTeamRequest,
|
||||
OrgMember,
|
||||
PatchTeamRequest,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
SpecialManagementEndpointEnums,
|
||||
|
|
@ -1956,6 +1957,7 @@ async def update_team(
|
|||
)
|
||||
async def patch_team(
|
||||
team_id: str,
|
||||
data: PatchTeamRequest,
|
||||
http_request: Request,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
litellm_changed_by: Annotated[
|
||||
|
|
@ -1968,11 +1970,12 @@ async def patch_team(
|
|||
"""
|
||||
Partially update a team using RFC 7386 JSON Merge Patch semantics.
|
||||
|
||||
`team_id` is taken from the path. `metadata` is merged with the team's stored
|
||||
metadata rather than replacing it: an omitted key is preserved, `key: null`
|
||||
deletes it, and any other value overwrites (recursing into nested objects).
|
||||
Every other field behaves exactly like `POST /team/update` (omitted preserves,
|
||||
a value overwrites). Returns the full updated team.
|
||||
`team_id` is taken from the path; a `team_id` in the body is accepted only when it
|
||||
matches. `metadata` is merged with the team's stored metadata rather than replacing
|
||||
it: an omitted key is preserved, `key: null` deletes it, and any other value
|
||||
overwrites (recursing into nested objects). Every other field behaves exactly like
|
||||
`POST /team/update` (omitted preserves, a value overwrites). Returns the full
|
||||
updated team.
|
||||
|
||||
```
|
||||
curl --location --request PATCH 'http://0.0.0.0:4000/team/8d916b1c-510d-4894-a334-1c16a93344f5' \
|
||||
|
|
@ -1992,21 +1995,15 @@ async def patch_team(
|
|||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
try:
|
||||
body = await http_request.json()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"})
|
||||
if not isinstance(body, dict):
|
||||
raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"})
|
||||
|
||||
body_team_id = body.pop("team_id", None)
|
||||
if body_team_id is not None and body_team_id != team_id:
|
||||
if data.team_id is not None and data.team_id != team_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"team_id in body ({body_team_id}) does not match team_id in path ({team_id})"},
|
||||
detail={"error": f"team_id in body ({data.team_id}) does not match team_id in path ({team_id})"},
|
||||
)
|
||||
|
||||
if "metadata" in body:
|
||||
patch_fields = data.model_dump(exclude_unset=True, exclude={"team_id"})
|
||||
|
||||
if "metadata" in patch_fields:
|
||||
existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
if existing_team_row is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -2014,9 +2011,9 @@ async def patch_team(
|
|||
detail={"error": f"Team not found, passed team_id={team_id}"},
|
||||
)
|
||||
existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {}
|
||||
body["metadata"] = apply_json_merge_patch(existing_metadata, body["metadata"])
|
||||
patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"])
|
||||
|
||||
update_request = UpdateTeamRequest(team_id=team_id, **body)
|
||||
update_request = UpdateTeamRequest(team_id=team_id, **patch_fields)
|
||||
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
|
|
@ -2375,7 +2372,15 @@ async def _add_team_members_to_team(
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: str,
|
||||
) -> Tuple[LiteLLM_TeamTable, List[LiteLLM_UserTable], List[LiteLLM_TeamMembership]]:
|
||||
"""Add team members to the team."""
|
||||
"""Add team members to the team.
|
||||
|
||||
The members_with_roles reconciliation runs inside a transaction that locks
|
||||
the team row with ``SELECT ... FOR UPDATE`` before reading the current
|
||||
membership. Concurrent /team/member_add calls for the same team therefore
|
||||
serialize on the row lock and each appends onto the other's committed
|
||||
result, instead of both rewriting the whole JSON array from a stale
|
||||
snapshot (which silently drops one member on the losing write).
|
||||
"""
|
||||
# Process and add new members
|
||||
updated_users, updated_team_memberships = await _process_team_members(
|
||||
data=data,
|
||||
|
|
@ -2385,19 +2390,22 @@ async def _add_team_members_to_team(
|
|||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
)
|
||||
|
||||
# Update team members list
|
||||
await _update_team_members_list(
|
||||
data=data,
|
||||
complete_team_data=complete_team_data,
|
||||
updated_users=updated_users,
|
||||
)
|
||||
async with prisma_client.tx() as tx:
|
||||
complete_team_data.members_with_roles = await TeamRepository(prisma_client).get_members_with_roles_locked(
|
||||
tx, data.team_id
|
||||
)
|
||||
|
||||
# ADD MEMBER TO TEAM
|
||||
_db_team_members = [m.model_dump() for m in complete_team_data.members_with_roles]
|
||||
updated_team = await TeamRepository(prisma_client).table.update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore
|
||||
)
|
||||
await _update_team_members_list(
|
||||
data=data,
|
||||
complete_team_data=complete_team_data,
|
||||
updated_users=updated_users,
|
||||
)
|
||||
|
||||
_db_team_members = [m.model_dump() for m in complete_team_data.members_with_roles]
|
||||
updated_team = await tx.litellm_teamtable.update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"members_with_roles": json.dumps(_db_team_members)},
|
||||
)
|
||||
|
||||
return updated_team, updated_users, updated_team_memberships
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import asyncio
|
|||
import base64
|
||||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
|
|
@ -62,6 +63,11 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
SSOIdentityAssertion,
|
||||
assertion_from_sso_login,
|
||||
retain_sso_identity_assertion_for_ema,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -253,11 +259,20 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic
|
|||
raise HTTPException(status_code=400, detail="Invalid CLI login session id")
|
||||
|
||||
cache_key = _get_cli_sso_flow_cache_key(cast(str, login_id))
|
||||
flow = cache.get_cache(key=cache_key)
|
||||
redis_cache = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
flow = redis_cache.get_cache(key=cache_key)
|
||||
else:
|
||||
flow = cache.get_cache(key=cache_key)
|
||||
if isinstance(flow, str):
|
||||
try:
|
||||
flow = json.loads(flow)
|
||||
except ValueError:
|
||||
flow = None
|
||||
if not isinstance(flow, dict) or "poll_secret_hash" not in flow:
|
||||
verbose_proxy_logger.warning(
|
||||
"CLI SSO login session not found in cache for login_id=%s. If the proxy runs multiple replicas, "
|
||||
"a shared Redis cache (enable_redis_auth_cache: true) is required for CLI login to work.",
|
||||
"a shared Redis cache is required for CLI login to work.",
|
||||
login_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
|
|
@ -265,7 +280,7 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic
|
|||
detail=(
|
||||
"CLI login session not found or expired. Run `litellm-proxy login` again. "
|
||||
"If this happens immediately after starting a login, the proxy is likely running multiple "
|
||||
"replicas without a shared cache; configure Redis with `enable_redis_auth_cache: true` "
|
||||
"replicas without a shared cache; configure a Redis cache "
|
||||
"so every replica can see the login session."
|
||||
),
|
||||
)
|
||||
|
|
@ -273,11 +288,12 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic
|
|||
|
||||
|
||||
def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
|
||||
cache.set_cache(
|
||||
key=_get_cli_sso_flow_cache_key(login_id),
|
||||
value=flow,
|
||||
ttl=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
)
|
||||
cache_key = _get_cli_sso_flow_cache_key(login_id)
|
||||
redis_cache = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
redis_cache.set_cache(key=cache_key, value=json.dumps(flow), ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
else:
|
||||
cache.set_cache(key=cache_key, value=flow, ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
|
||||
|
||||
def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool:
|
||||
|
|
@ -588,11 +604,11 @@ def _render_cli_sso_verification_page(
|
|||
|
||||
@router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False)
|
||||
async def cli_sso_start(request: Request):
|
||||
from litellm.proxy.proxy_server import general_settings, user_api_key_cache
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache, general_settings
|
||||
|
||||
_check_cli_sso_start_rate_limit(
|
||||
request=request,
|
||||
cache=user_api_key_cache,
|
||||
cache=cli_sso_session_cache,
|
||||
use_x_forwarded_for=bool((general_settings or {}).get("use_x_forwarded_for", False)),
|
||||
)
|
||||
|
||||
|
|
@ -607,7 +623,7 @@ async def cli_sso_start(request: Request):
|
|||
"user_code_verified": False,
|
||||
"session_data": None,
|
||||
}
|
||||
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
verification_uri_complete: str | None = (
|
||||
(
|
||||
|
|
@ -639,9 +655,9 @@ async def cli_sso_complete(request: Request, login_id: str):
|
|||
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
|
||||
render_cli_sso_success_page,
|
||||
)
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache
|
||||
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cli_sso_session_cache)
|
||||
if not flow.get("sso_complete") or not flow.get("session_data"):
|
||||
raise HTTPException(status_code=400, detail="CLI login is not ready")
|
||||
|
||||
|
|
@ -665,7 +681,7 @@ async def cli_sso_complete(request: Request, login_id: str):
|
|||
raise HTTPException(status_code=400, detail="Invalid verification code")
|
||||
|
||||
flow["user_code_verified"] = True
|
||||
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
html_content = render_cli_sso_success_page()
|
||||
return HTMLResponse(content=html_content, status_code=200)
|
||||
|
|
@ -856,10 +872,10 @@ async def google_login(
|
|||
Example:
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
cli_sso_session_cache,
|
||||
general_settings,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
user_custom_ui_sso_sign_in_handler,
|
||||
)
|
||||
|
||||
|
|
@ -907,7 +923,7 @@ async def google_login(
|
|||
)
|
||||
|
||||
if source == LITELLM_CLI_SOURCE_IDENTIFIER:
|
||||
_get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
|
||||
_get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
|
||||
|
||||
# Store CLI login handle in state for OAuth flow
|
||||
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
|
||||
|
|
@ -1311,12 +1327,15 @@ async def get_generic_sso_response(
|
|||
sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control
|
||||
generic_client_id: str,
|
||||
redirect_url: str,
|
||||
) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload)
|
||||
) -> tuple[
|
||||
Union[OpenID, dict], dict | None, dict | None, SSOIdentityAssertion | None
|
||||
]: # (result, received_response, access_token_payload, sso_assertion)
|
||||
# make generic sso provider
|
||||
from fastapi_sso.sso.base import DiscoveryDocument
|
||||
from fastapi_sso.sso.generic import create_provider
|
||||
|
||||
received_response: Optional[dict] = None
|
||||
sso_assertion: SSOIdentityAssertion | None = None
|
||||
|
||||
# Setup environment variables
|
||||
(
|
||||
|
|
@ -1450,6 +1469,9 @@ async def get_generic_sso_response(
|
|||
# Assign directly rather than relying on nonlocal mutation so that Pyright
|
||||
# can track that received_response is non-None from this point on.
|
||||
received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS}
|
||||
sso_assertion = assertion_from_sso_login(
|
||||
combined_response.get("id_token"), combined_response.get("refresh_token")
|
||||
)
|
||||
# In the PKCE path verify_and_process is skipped, so generic_sso.access_token
|
||||
# is never set. Read the token directly from the exchange response instead so
|
||||
# process_sso_jwt_access_token can extract JWT-embedded roles/teams.
|
||||
|
|
@ -1461,6 +1483,7 @@ async def get_generic_sso_response(
|
|||
headers=additional_generic_sso_headers_dict,
|
||||
)
|
||||
access_token_str = generic_sso.access_token
|
||||
sso_assertion = assertion_from_sso_login(generic_sso.id_token, generic_sso.refresh_token)
|
||||
|
||||
access_token_payload = process_sso_jwt_access_token(
|
||||
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
|
||||
|
|
@ -1480,7 +1503,7 @@ async def get_generic_sso_response(
|
|||
additional_generic_sso_headers_dict,
|
||||
)
|
||||
verbose_proxy_logger.debug("generic result: %s", result)
|
||||
return result or {}, received_response, access_token_payload
|
||||
return result or {}, received_response, access_token_payload, sso_assertion
|
||||
|
||||
|
||||
async def create_team_member_add_task(team_id, user_info):
|
||||
|
|
@ -1812,6 +1835,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
generic_client_id = os.getenv("GENERIC_CLIENT_ID", None)
|
||||
received_response: Optional[dict] = None
|
||||
access_token_payload: Optional[dict] = None
|
||||
sso_assertion: SSOIdentityAssertion | None = None
|
||||
# get url from request
|
||||
if master_key is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -1842,6 +1866,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
result,
|
||||
received_response,
|
||||
access_token_payload,
|
||||
sso_assertion,
|
||||
) = await get_generic_sso_response(
|
||||
request=request,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
@ -1869,6 +1894,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
prefill_user_code=prefill_user_code,
|
||||
result=result,
|
||||
received_response=received_response,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
|
||||
# Control-plane cross-origin: read return_to from cookie.
|
||||
|
|
@ -1884,6 +1910,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
access_token_payload=access_token_payload,
|
||||
jwt_handler=jwt_handler,
|
||||
return_to=cp_return_to,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1941,8 +1968,10 @@ async def _complete_cli_sso_callback_session(
|
|||
user_defined_values: Optional[SSOUserDefinedValues],
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
cli_sso_session_cache: DualCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prefill_user_code: str | None = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
):
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
|
|
@ -1962,6 +1991,8 @@ async def _complete_cli_sso_callback_session(
|
|||
if not user_info.user_id:
|
||||
raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
|
||||
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
|
||||
|
||||
teams: List[str] = []
|
||||
if hasattr(user_info, "teams") and user_info.teams:
|
||||
teams = user_info.teams if isinstance(user_info.teams, list) else []
|
||||
|
|
@ -1987,7 +2018,7 @@ async def _complete_cli_sso_callback_session(
|
|||
flow["sso_complete"] = True
|
||||
browser_complete_token = secrets.token_urlsafe(32)
|
||||
flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token)
|
||||
_set_cli_sso_flow(login_id=key, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=key, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}"
|
||||
|
|
@ -2012,18 +2043,20 @@ async def cli_sso_callback(
|
|||
result: Optional[Union[OpenID, dict]] = None,
|
||||
received_response: Optional[dict] = None,
|
||||
prefill_user_code: str | None = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
):
|
||||
"""CLI SSO callback - stores session info for JWT generation on polling"""
|
||||
verbose_proxy_logger.info("CLI SSO callback")
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
cli_sso_session_cache,
|
||||
general_settings,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
|
@ -2063,8 +2096,10 @@ async def cli_sso_callback(
|
|||
user_defined_values=user_defined_values,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
cli_sso_session_cache=cli_sso_session_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prefill_user_code=prefill_user_code,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
except ProxyException:
|
||||
raise
|
||||
|
|
@ -2093,10 +2128,10 @@ async def cli_poll_key(
|
|||
team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache
|
||||
|
||||
try:
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=cli_sso_session_cache)
|
||||
if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret):
|
||||
raise HTTPException(status_code=403, detail="Invalid CLI polling secret")
|
||||
|
||||
|
|
@ -2171,7 +2206,7 @@ async def cli_poll_key(
|
|||
)
|
||||
|
||||
# Delete cache entry (single-use)
|
||||
user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
|
||||
cli_sso_session_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
|
||||
|
||||
verbose_proxy_logger.info(f"CLI JWT generated for user: {user_id}, team: {team_id}")
|
||||
poll_response = {
|
||||
|
|
@ -3018,6 +3053,7 @@ class SSOAuthenticationHandler:
|
|||
access_token_payload: Optional[dict] = None,
|
||||
jwt_handler: Optional[JWTHandler] = None,
|
||||
return_to: Optional[str] = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
) -> RedirectResponse:
|
||||
import jwt
|
||||
|
||||
|
|
@ -3148,6 +3184,9 @@ class SSOAuthenticationHandler:
|
|||
},
|
||||
)
|
||||
|
||||
if isinstance(user_id, str) and user_id:
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
|
||||
|
||||
disabled_non_admin_personal_key_creation = get_disabled_non_admin_personal_key_creation()
|
||||
litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
|
||||
|
||||
|
|
@ -4241,6 +4280,7 @@ async def debug_sso_callback(request: Request):
|
|||
result,
|
||||
received_response,
|
||||
access_token_payload,
|
||||
_sso_assertion,
|
||||
) = await get_generic_sso_response(
|
||||
request=request,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ organizations, teams, and keys.
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Set, Union
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
|
||||
|
|
@ -64,6 +65,57 @@ async def attach_object_permission_to_dict(
|
|||
return data_dict
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ObjectPermissionUpsert:
|
||||
object_permission_id: str
|
||||
record: dict[str, object]
|
||||
|
||||
|
||||
async def prepare_object_permission_upsert(
|
||||
new_object_permission: Mapping[str, object],
|
||||
existing_object_permission_id: str | None,
|
||||
prisma_client: PrismaClient,
|
||||
) -> ObjectPermissionUpsert:
|
||||
"""
|
||||
Read-and-merge half of an object permission upsert; performs no writes.
|
||||
|
||||
Merges the sent grants over the existing row (looked up by
|
||||
``existing_object_permission_id``, or a fresh uuid when the entity has none) and
|
||||
returns the id plus the full record to upsert. The id is pinned inside the record
|
||||
because the column has ``@default(uuid())``, so a create without it would mint a
|
||||
different id than the one the caller links. ``mcp_tool_permissions`` is serialized
|
||||
to a JSON string to avoid GraphQL parsing issues (e.g. server IDs starting with
|
||||
"3e64" being interpreted as floats).
|
||||
|
||||
Keeping this separate from the write lets callers run the upsert inside the same
|
||||
transaction as the row that links ``object_permission_id``, so a rolled-back
|
||||
update cannot leave permission changes live.
|
||||
"""
|
||||
object_permission_id = existing_object_permission_id or str(uuid.uuid4())
|
||||
existing_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
)
|
||||
existing_fields: dict[str, object] = (
|
||||
existing_object_permission.model_dump(exclude_unset=True, exclude_none=True)
|
||||
if existing_object_permission is not None
|
||||
else {}
|
||||
)
|
||||
merged: dict[str, object] = {
|
||||
**existing_fields,
|
||||
**new_object_permission,
|
||||
"object_permission_id": object_permission_id,
|
||||
}
|
||||
record: dict[str, object] = {
|
||||
**merged,
|
||||
**(
|
||||
{"mcp_tool_permissions": safe_dumps(merged["mcp_tool_permissions"])}
|
||||
if "mcp_tool_permissions" in merged
|
||||
else {}
|
||||
),
|
||||
}
|
||||
return ObjectPermissionUpsert(object_permission_id=object_permission_id, record=record)
|
||||
|
||||
|
||||
async def handle_update_object_permission_common(
|
||||
data_json: Dict,
|
||||
existing_object_permission_id: Optional[str],
|
||||
|
|
@ -93,50 +145,23 @@ async def handle_update_object_permission_common(
|
|||
if prisma_client is None:
|
||||
raise ValueError("Prisma client not found")
|
||||
|
||||
#########################################################
|
||||
# Ensure `object_permission` is not added to the data_json
|
||||
# We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission: Union[dict, str] = data_json.pop("object_permission", None)
|
||||
new_object_permission: Union[dict, str, None] = data_json.pop("object_permission", None)
|
||||
if new_object_permission is None:
|
||||
return None
|
||||
|
||||
# Lookup existing object permission ID and update that entry
|
||||
object_permission_id_to_use: str = existing_object_permission_id or str(uuid.uuid4())
|
||||
existing_object_permissions_dict: Dict = {}
|
||||
|
||||
existing_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique(
|
||||
where={"object_permission_id": object_permission_id_to_use},
|
||||
)
|
||||
|
||||
# Update the object permission
|
||||
if existing_object_permission is not None:
|
||||
existing_object_permissions_dict = existing_object_permission.model_dump(exclude_unset=True, exclude_none=True)
|
||||
|
||||
# Handle string JSON object permission
|
||||
if isinstance(new_object_permission, str):
|
||||
new_object_permission = json.loads(new_object_permission)
|
||||
|
||||
if isinstance(new_object_permission, dict):
|
||||
existing_object_permissions_dict.update(new_object_permission)
|
||||
|
||||
#########################################################
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
#########################################################
|
||||
if "mcp_tool_permissions" in existing_object_permissions_dict:
|
||||
existing_object_permissions_dict["mcp_tool_permissions"] = safe_dumps(
|
||||
existing_object_permissions_dict["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
upsert = await prepare_object_permission_upsert(
|
||||
new_object_permission=new_object_permission if isinstance(new_object_permission, dict) else {},
|
||||
existing_object_permission_id=existing_object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
created_object_permission_row = await ObjectPermissionRepository(prisma_client).table.upsert(
|
||||
where={"object_permission_id": object_permission_id_to_use},
|
||||
where={"object_permission_id": upsert.object_permission_id},
|
||||
data={
|
||||
"create": existing_object_permissions_dict,
|
||||
"update": existing_object_permissions_dict,
|
||||
"create": upsert.record,
|
||||
"update": upsert.record,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -252,6 +252,21 @@ async def _resolve_member_budget_id(
|
|||
return response.budget_id
|
||||
|
||||
|
||||
async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, team_id: str) -> None:
|
||||
"""Append team_id to a user's teams array, only if it is not already present.
|
||||
|
||||
The row-level filter makes the append a no-op once the team is present, so
|
||||
repeated or concurrent adds of the same team cannot accumulate duplicate
|
||||
team ids in user.teams (a duplicate also breaks auth logic that keys off the
|
||||
number of teams a user belongs to). Teams added concurrently for a different
|
||||
team id are unaffected, since each update filters on its own team id.
|
||||
"""
|
||||
await UserRepository(prisma_client).table.update_many(
|
||||
where={"user_id": user_id, "NOT": {"teams": {"has": team_id}}},
|
||||
data={"teams": {"push": [team_id]}},
|
||||
)
|
||||
|
||||
|
||||
async def add_new_member(
|
||||
new_member: Member,
|
||||
max_budget_in_team: Optional[float],
|
||||
|
|
@ -276,13 +291,16 @@ async def add_new_member(
|
|||
## ADD TEAM ID, to USER TABLE IF NEW ##
|
||||
if new_member.user_id is not None:
|
||||
new_user_defaults = get_new_internal_user_defaults(user_id=new_member.user_id)
|
||||
# Upsert ensures the user row exists atomically (no create race when the
|
||||
# same new user is provisioned concurrently), seeding teams on create.
|
||||
# The teams append lives in the filtered update below rather than the
|
||||
# upsert's update branch so an already-existing user does not get a
|
||||
# duplicate team id.
|
||||
_returned_user = await UserRepository(prisma_client).table.upsert(
|
||||
where={"user_id": new_member.user_id},
|
||||
data={
|
||||
"update": {"teams": {"push": [team_id]}},
|
||||
"create": {"teams": [team_id], **new_user_defaults}, # type: ignore
|
||||
},
|
||||
data={"create": {"teams": [team_id], **new_user_defaults}, "update": {}},
|
||||
)
|
||||
await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id)
|
||||
if _returned_user is not None:
|
||||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
elif new_member.user_email is not None:
|
||||
|
|
@ -302,12 +320,8 @@ async def add_new_member(
|
|||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
elif len(existing_user_row) == 1:
|
||||
user_info = existing_user_row[0]
|
||||
_returned_user = await UserRepository(prisma_client).table.update(
|
||||
where={"user_id": user_info.user_id}, # type: ignore
|
||||
data={"teams": {"push": [team_id]}},
|
||||
)
|
||||
if _returned_user is not None:
|
||||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id)
|
||||
returned_user = LiteLLM_UserTable(**user_info.model_dump())
|
||||
elif len(existing_user_row) > 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
|
|||
|
|
@ -198,13 +198,9 @@
|
|||
"icon_url": "https://cdn.simpleicons.org/googledrive",
|
||||
"category": "Productivity",
|
||||
"registry_url": null,
|
||||
"transport": "stdio",
|
||||
"command": "npx",
|
||||
"args": ["-y", "@modelcontextprotocol/server-gdrive"],
|
||||
"env_vars": [
|
||||
{"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false},
|
||||
{"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true}
|
||||
]
|
||||
"transport": "http",
|
||||
"url": "https://drivemcp.googleapis.com/mcp/v1",
|
||||
"env_vars": []
|
||||
},
|
||||
{
|
||||
"name": "google_calendar",
|
||||
|
|
|
|||
|
|
@ -226,6 +226,7 @@ from litellm.constants import (
|
|||
APSCHEDULER_MAX_INSTANCES,
|
||||
APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
APSCHEDULER_REPLACE_EXISTING,
|
||||
CLI_SSO_SESSION_TTL_SECONDS,
|
||||
DAYS_IN_A_MONTH,
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
|
|
@ -236,6 +237,7 @@ from litellm.constants import (
|
|||
PROXY_BATCH_WRITE_AT,
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
|
||||
)
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
|
@ -302,6 +304,11 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.config_resolvers import resolve_fields
|
||||
from litellm.proxy.config_resolvers.alerting import (
|
||||
EMAIL_DESCRIPTORS,
|
||||
SLACK_DESCRIPTORS,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
|
|
@ -319,7 +326,10 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
from litellm.proxy.common_utils.proxy_state import ProxyState
|
||||
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
|
||||
from litellm.proxy.common_utils.swagger_utils import ERROR_RESPONSES
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
get_budget_reset_settings,
|
||||
get_budget_reset_time,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
get_management_object_ttl,
|
||||
|
|
@ -1154,9 +1164,9 @@ _OPENAPI_HTTP_METHODS = {
|
|||
# Credentials surfaced by `/get/config/callbacks` in the alerting block: the
|
||||
# full Slack incoming-webhook URL is itself a credential, and the SMTP
|
||||
# password is a service password. Masked on read so plaintext never reaches
|
||||
# the UI. Kept here at module scope to match the analogous
|
||||
# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO
|
||||
# and cache endpoint files.
|
||||
# the UI. Kept here at module scope to match the analogous descriptor
|
||||
# `is_secret` flags in litellm.proxy.config_resolvers and the
|
||||
# `_CACHE_SENSITIVE_FIELDS` constant in the cache endpoint file.
|
||||
_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
|
||||
|
||||
|
||||
|
|
@ -1966,6 +1976,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache(
|
|||
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
|
||||
)
|
||||
spend_counter_cache = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value)
|
||||
cli_sso_session_cache = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=user_api_key_cache)
|
||||
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
|
||||
redis_usage_cache: Optional[RedisCache] = None # redis cache used for tracking spend, tpm/rpm limits
|
||||
|
|
@ -1998,6 +2009,7 @@ proxy_budget_rescheduler_min_time = PROXY_BUDGET_RESCHEDULER_MIN_TIME
|
|||
proxy_budget_rescheduler_max_time = PROXY_BUDGET_RESCHEDULER_MAX_TIME
|
||||
proxy_batch_polling_interval = PROXY_BATCH_POLLING_INTERVAL
|
||||
proxy_batch_write_at = PROXY_BATCH_WRITE_AT
|
||||
proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS
|
||||
litellm_master_key_hash = None
|
||||
disable_spend_logs = False
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -3691,13 +3703,22 @@ def _build_redis_usage_cache_from_environment() -> RedisCache | None:
|
|||
def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: bool) -> None:
|
||||
"""
|
||||
Wires an established coordination Redis into the proxy-level caches that
|
||||
consume it directly: the spend counter cache, the cluster-wide config
|
||||
cache, and (only when opted in) the virtual-key auth cache.
|
||||
consume it directly: the spend counter cache, the CLI SSO login-session
|
||||
cache, the cluster-wide config cache, and (only when opted in) the
|
||||
virtual-key auth cache.
|
||||
|
||||
The CLI SSO login-session cache is always backed by Redis when available so
|
||||
that the browser SSO flow behind `lite login` survives landing on different
|
||||
workers; it must not be gated behind enable_redis_auth_cache.
|
||||
"""
|
||||
spend_counter_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
default_redis_ttl=litellm.default_redis_ttl,
|
||||
)
|
||||
cli_sso_session_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
)
|
||||
if enable_redis_auth_cache is True:
|
||||
user_api_key_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
|
|
@ -3879,7 +3900,7 @@ class ProxyConfig:
|
|||
del config["include"]
|
||||
return config
|
||||
|
||||
async def save_config(self, new_config: dict):
|
||||
async def save_config(self, new_config: dict, include_env_vars: bool = False):
|
||||
global prisma_client, general_settings, user_config_file_path, store_model_in_db
|
||||
# Load existing config
|
||||
## DB - writes valid config to db
|
||||
|
|
@ -3896,6 +3917,17 @@ class ProxyConfig:
|
|||
# Make a copy to avoid mutating the original config
|
||||
config_to_save = new_config.copy()
|
||||
|
||||
# environment_variables are persisted to the DB only when a caller
|
||||
# explicitly opts in. Most callers reach save_config after
|
||||
# get_config() merged YAML + OS env into new_config (with
|
||||
# os.environ/ placeholders already resolved to plaintext), so
|
||||
# persisting them here would snapshot file/container env vars into
|
||||
# a config row that then shadows those sources on every restart.
|
||||
# The dedicated /config/update path writes env vars directly, so
|
||||
# no current caller needs include_env_vars=True.
|
||||
if not include_env_vars:
|
||||
config_to_save.pop("environment_variables", None)
|
||||
|
||||
# SECURITY: Always encrypt environment_variables before DB write.
|
||||
# _encrypt_env_variables_for_db is idempotent — a caller that
|
||||
# already encrypted the values (or re-submitted ciphertext read
|
||||
|
|
@ -3913,6 +3945,38 @@ class ProxyConfig:
|
|||
with open(f"{user_config_file_path}", "w") as config_file:
|
||||
yaml.dump(new_config, config_file, default_flow_style=False)
|
||||
|
||||
async def save_environment_variables(self, updates: dict[str, str | None]) -> None:
|
||||
"""Persist specific environment variables to the DB config row.
|
||||
|
||||
Each key in ``updates`` is written to the ``environment_variables``
|
||||
config row; a ``None`` value deletes that key. Env vars the caller does
|
||||
not name are preserved, so a caller that owns a couple of keys can
|
||||
update just those without snapshotting unrelated (YAML/OS-sourced)
|
||||
values the way a full ``save_config`` write would. No-op when config is
|
||||
not DB-backed.
|
||||
"""
|
||||
global prisma_client, general_settings, store_model_in_db
|
||||
if prisma_client is None or not (general_settings.get("store_model_in_db", False) is True or store_model_in_db):
|
||||
return
|
||||
|
||||
row = await ConfigRepository(prisma_client).table.find_first(where={"param_name": "environment_variables"})
|
||||
existing: dict = dict(row.param_value) if row is not None and row.param_value is not None else {}
|
||||
|
||||
to_set = {k: v for k, v in updates.items() if v is not None}
|
||||
encrypted = self._encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {}
|
||||
deleted_keys = {k for k, v in updates.items() if v is None}
|
||||
merged = {**{k: v for k, v in existing.items() if k not in deleted_keys}, **encrypted}
|
||||
|
||||
serialized = json.dumps(merged)
|
||||
await ConfigRepository(prisma_client).table.upsert(
|
||||
where={"param_name": "environment_variables"},
|
||||
data={
|
||||
"create": {"param_name": "environment_variables", "param_value": serialized},
|
||||
"update": {"param_value": serialized},
|
||||
},
|
||||
)
|
||||
await invalidate_config_param("environment_variables")
|
||||
|
||||
def _check_for_os_environ_vars(
|
||||
self, config: dict, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH
|
||||
) -> dict:
|
||||
|
|
@ -4291,6 +4355,7 @@ class ProxyConfig:
|
|||
open_telemetry_logger, \
|
||||
health_check_details, \
|
||||
proxy_batch_polling_interval, \
|
||||
proxy_config_reload_interval_seconds, \
|
||||
config_passthrough_endpoints
|
||||
|
||||
config: dict = await self.get_config(config_file_path=config_file_path)
|
||||
|
|
@ -4569,6 +4634,11 @@ class ProxyConfig:
|
|||
verbose_proxy_logger.debug(
|
||||
f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, native_background_mode={native_background_mode}, ttl={polling_cache_ttl}{reset_color_code}"
|
||||
)
|
||||
elif key == "max_ui_session_budget":
|
||||
litellm.max_ui_session_budget = float(value) if value is not None else None
|
||||
verbose_proxy_logger.debug(
|
||||
f"{blue_color_code} setting litellm.max_ui_session_budget={litellm.max_ui_session_budget}{reset_color_code}"
|
||||
)
|
||||
elif key == "default_team_settings":
|
||||
for idx, team_setting in enumerate(value): # run through pydantic validation
|
||||
try:
|
||||
|
|
@ -4597,6 +4667,13 @@ class ProxyConfig:
|
|||
litellm.json_logs = True
|
||||
litellm._turn_on_json()
|
||||
verbose_proxy_logger.debug(f"{blue_color_code} Enabled JSON logging via config{reset_color_code}")
|
||||
elif key == "budget_reset_time":
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
parse_budget_reset_time,
|
||||
)
|
||||
|
||||
parse_budget_reset_time(value)
|
||||
setattr(litellm, key, value)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
f"{blue_color_code} setting litellm.{key}={_redact_general_setting_value(key, value, is_full_admin=False)}{reset_color_code}"
|
||||
|
|
@ -4773,6 +4850,10 @@ class ProxyConfig:
|
|||
)
|
||||
## BATCH WRITER ##
|
||||
proxy_batch_write_at = general_settings.get("proxy_batch_write_at", proxy_batch_write_at)
|
||||
## DB CONFIG RELOAD INTERVAL ##
|
||||
proxy_config_reload_interval_seconds = general_settings.get(
|
||||
"proxy_config_reload_interval_seconds", proxy_config_reload_interval_seconds
|
||||
)
|
||||
## DISABLE SPEND LOGS ## - gives a perf improvement
|
||||
disable_spend_logs = general_settings.get("disable_spend_logs", disable_spend_logs)
|
||||
### BACKGROUND HEALTH CHECKS ###
|
||||
|
|
@ -7868,6 +7949,7 @@ class ProxyStartupEvent:
|
|||
budget_reset_job = ResetBudgetJob(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prisma_client=prisma_client,
|
||||
reset_settings=get_budget_reset_settings(),
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
|
|
@ -7944,12 +8026,20 @@ class ProxyStartupEvent:
|
|||
verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e))
|
||||
|
||||
if store_model_in_db is True:
|
||||
config_reload_interval_seconds = proxy_config_reload_interval_seconds
|
||||
if not isinstance(config_reload_interval_seconds, int) or config_reload_interval_seconds <= 0:
|
||||
verbose_proxy_logger.warning(
|
||||
"proxy_config_reload_interval_seconds=%s must be a positive integer; falling back to 30s",
|
||||
config_reload_interval_seconds,
|
||||
)
|
||||
config_reload_interval_seconds = 30
|
||||
|
||||
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
|
||||
# Frequent polling was causing excessive memory allocations
|
||||
scheduler.add_job(
|
||||
proxy_config.add_deployment,
|
||||
"interval",
|
||||
seconds=30, # increased from 10s to reduce memory pressure
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client, proxy_logging_obj],
|
||||
id="add_deployment_job",
|
||||
|
|
@ -7964,7 +8054,7 @@ class ProxyStartupEvent:
|
|||
scheduler.add_job(
|
||||
proxy_config.get_credentials,
|
||||
"interval",
|
||||
seconds=30, # increased from 10s to reduce memory pressure
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client],
|
||||
id="get_credentials_job",
|
||||
|
|
@ -14845,10 +14935,11 @@ GeneralSettingsUILiteLLMValue = Union[float, bool, str, None]
|
|||
|
||||
|
||||
class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
|
||||
type: Literal["Float", "Boolean", "Select"]
|
||||
type: Literal["Float", "Dollar", "Boolean", "Select"]
|
||||
description: str
|
||||
options: NotRequired[tuple[str, ...]]
|
||||
tab: NotRequired[str] # Admin UI sub-tab this field renders under; None groups it with the rest
|
||||
default: NotRequired[float] # reset/clear restores this instead of None; fields whose None means fail-open set it
|
||||
|
||||
|
||||
_GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec] = {
|
||||
|
|
@ -14874,21 +14965,32 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec
|
|||
"tab": "prompt_caching",
|
||||
"description": "Empty uses Anthropic's 5m default. 1h suits long sessions but doubles the cache write cost.",
|
||||
},
|
||||
"max_ui_session_budget": {
|
||||
"type": "Dollar",
|
||||
"default": 1.0,
|
||||
"description": (
|
||||
"USD spend cap for each dashboard login session; covers LLM calls made from the dashboard "
|
||||
"such as the playground and auto router Test Connection. Each login starts a fresh session "
|
||||
"with this budget. Clearing restores the $1 default."
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _general_settings_ui_litellm_default(
|
||||
field_type: Literal["Float", "Boolean", "Select"],
|
||||
spec: GeneralSettingsUILiteLLMFieldSpec,
|
||||
) -> GeneralSettingsUILiteLLMValue:
|
||||
"""The value a field falls back to when it is cleared or reset."""
|
||||
return False if field_type == "Boolean" else None
|
||||
if "default" in spec:
|
||||
return spec["default"]
|
||||
return False if spec["type"] == "Boolean" else None
|
||||
|
||||
|
||||
def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> GeneralSettingsUILiteLLMValue:
|
||||
spec = _GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]
|
||||
field_type = spec["type"]
|
||||
if value is None or value == "":
|
||||
return _general_settings_ui_litellm_default(field_type)
|
||||
return _general_settings_ui_litellm_default(spec)
|
||||
match field_type:
|
||||
case "Boolean":
|
||||
if not isinstance(value, bool):
|
||||
|
|
@ -14912,6 +15014,13 @@ def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) ->
|
|||
detail={"error": f"{field_name} must be a number in (0, 1] or empty"},
|
||||
)
|
||||
return float(value)
|
||||
case "Dollar":
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)) or float(value) <= 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"{field_name} must be a positive dollar amount or empty"},
|
||||
)
|
||||
return float(value)
|
||||
case _:
|
||||
assert_never(field_type)
|
||||
|
||||
|
|
@ -14934,7 +15043,7 @@ async def _persist_general_settings_ui_litellm_field(
|
|||
async def _reset_general_settings_ui_litellm_field(field_name: str, user_api_key_dict: UserAPIKeyAuth) -> dict:
|
||||
config = await proxy_config.get_config()
|
||||
before_value = config.get("litellm_settings", {}).get(field_name)
|
||||
default_value = _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]["type"])
|
||||
default_value = _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name])
|
||||
setattr(litellm, field_name, default_value)
|
||||
if "litellm_settings" in config:
|
||||
config["litellm_settings"].pop(field_name, None)
|
||||
|
|
@ -14998,6 +15107,7 @@ async def get_config_list(
|
|||
"global_max_parallel_requests": {"type": "Integer"},
|
||||
"max_request_size_mb": {"type": "Integer"},
|
||||
"max_response_size_mb": {"type": "Integer"},
|
||||
"proxy_config_reload_interval_seconds": {"type": "Integer"},
|
||||
"pass_through_endpoints": {"type": "PydanticModel"},
|
||||
"store_model_in_db": {"type": "Boolean"},
|
||||
"store_prompts_in_spend_logs": {"type": "Boolean"},
|
||||
|
|
@ -15108,7 +15218,7 @@ async def get_config_list(
|
|||
)
|
||||
for litellm_field_name, spec in _GENERAL_SETTINGS_UI_LITELLM_FIELDS.items():
|
||||
current_value: GeneralSettingsUILiteLLMValue = getattr(litellm, litellm_field_name, None)
|
||||
default_value = _general_settings_ui_litellm_default(spec["type"])
|
||||
default_value = _general_settings_ui_litellm_default(spec)
|
||||
stored_in_db_litellm: Optional[bool]
|
||||
if litellm_field_name in db_litellm_settings:
|
||||
stored_in_db_litellm = True
|
||||
|
|
@ -15386,14 +15496,10 @@ async def get_config(
|
|||
_alerting = _general_settings.get("alerting", [])
|
||||
alerting_data = []
|
||||
if "slack" in _alerting:
|
||||
_slack_vars = [
|
||||
"SLACK_WEBHOOK_URL",
|
||||
]
|
||||
_slack_env_vars = {
|
||||
_var: (value if (value := environment_variables.get(_var)) is not None else os.getenv(_var))
|
||||
for _var in _slack_vars
|
||||
}
|
||||
_slack_env_vars = _apply_alerting_env_role_gate(_slack_env_vars, is_full_admin)
|
||||
_slack_values, _ = resolve_fields(
|
||||
SLACK_DESCRIPTORS, environment_variables, os.environ, empty_db_is_set=True
|
||||
)
|
||||
_slack_env_vars = _apply_alerting_env_role_gate(_slack_values, is_full_admin)
|
||||
|
||||
_alerting_types = proxy_logging_obj.slack_alerting_instance.alert_types
|
||||
_all_alert_types = proxy_logging_obj.slack_alerting_instance._all_possible_alert_types()
|
||||
|
|
@ -15409,19 +15515,8 @@ async def get_config(
|
|||
}
|
||||
)
|
||||
# pass email alerting vars
|
||||
_email_vars = [
|
||||
"SMTP_HOST",
|
||||
"SMTP_PORT",
|
||||
"SMTP_USERNAME",
|
||||
"SMTP_PASSWORD",
|
||||
"SMTP_SENDER_EMAIL",
|
||||
"TEST_EMAIL_ADDRESS",
|
||||
"EMAIL_LOGO_URL",
|
||||
"EMAIL_SUPPORT_CONTACT",
|
||||
]
|
||||
_email_env_vars = _apply_alerting_env_role_gate(
|
||||
{_var: environment_variables.get(_var) for _var in _email_vars}, is_full_admin
|
||||
)
|
||||
_email_values, _ = resolve_fields(EMAIL_DESCRIPTORS, environment_variables, os.environ, empty_db_is_set=True)
|
||||
_email_env_vars = _apply_alerting_env_role_gate(_email_values, is_full_admin)
|
||||
|
||||
alerting_data.append(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
#### CRUD ENDPOINTS for UI Settings #####
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple, Type, Union
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
|
@ -13,6 +15,11 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.config_resolvers.sso import (
|
||||
SSO_FIELD_ENV_VARS,
|
||||
SSO_SECRET_FIELDS,
|
||||
resolve_sso_config,
|
||||
)
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
SSOConfigRepository,
|
||||
|
|
@ -25,17 +32,45 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
|
|||
|
||||
router = APIRouter()
|
||||
|
||||
# SSO secret fields returned by /get/sso_settings. These are masked on read so
|
||||
# the UI can show "(set)" without ever transporting the plaintext OAuth secret
|
||||
# off the server, matching the write-once + masked-on-read contract used for
|
||||
# the HashiCorp Vault config override.
|
||||
_SSO_SENSITIVE_FIELDS: Set[str] = {
|
||||
"google_client_secret",
|
||||
"microsoft_client_secret",
|
||||
"generic_client_secret",
|
||||
# Maps each UIThemeConfig field to the env var the UI branding path reads it
|
||||
# from. /update/ui_theme_settings writes both the stored ui_theme_config and
|
||||
# these env vars, so /get/ui_theme_settings resolves the same env vars to
|
||||
# reflect a deployment branded purely through process env.
|
||||
_UI_THEME_FIELD_ENV_VARS: dict[str, str] = {
|
||||
"logo_url": "UI_LOGO_PATH",
|
||||
"favicon_url": "LITELLM_FAVICON_URL",
|
||||
}
|
||||
|
||||
|
||||
def _is_public_http_url(value: str | None) -> bool:
|
||||
"""Whether a value is a plain http(s) URL with a host, safe to disclose publicly."""
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
return False
|
||||
parsed = urlparse(value.strip())
|
||||
return parsed.scheme in ("http", "https") and bool(parsed.netloc)
|
||||
|
||||
|
||||
def _resolve_ui_theme_field(stored_values: Mapping[str, Any], field_name: str) -> str | None:
|
||||
"""Resolve one UI theme field to the value the branding path actually uses.
|
||||
|
||||
The stored ui_theme_config wins; a field absent or blank there falls back to
|
||||
the process environment. The branding path reads the env var, and stored
|
||||
settings reach it by being pushed into the environment on save, so a value
|
||||
supplied only as a process env var is live even though no stored entry exists.
|
||||
|
||||
This endpoint is unauthenticated, so the env fallback only surfaces a public
|
||||
http(s) URL: an operator can point UI_LOGO_PATH at a local filesystem path
|
||||
(the branding path serves it server-side), and that path must not be
|
||||
disclosed to anonymous callers. A stored value is already validated as a
|
||||
public URL on write, so it passes through.
|
||||
"""
|
||||
stored = stored_values.get(field_name)
|
||||
if isinstance(stored, str) and stored.strip():
|
||||
return stored
|
||||
env_value = os.environ.get(_UI_THEME_FIELD_ENV_VARS[field_name])
|
||||
return env_value if _is_public_http_url(env_value) else None
|
||||
|
||||
|
||||
class IPAddress(BaseModel):
|
||||
ip: str
|
||||
|
||||
|
|
@ -69,7 +104,8 @@ class SettingsResponse(BaseModel):
|
|||
class SSOSettingsResponse(SettingsResponse):
|
||||
"""Response model for SSO settings"""
|
||||
|
||||
pass
|
||||
provenance: Dict[str, str] = Field(default_factory=dict)
|
||||
"""Per-field source of each value: 'db', 'env', 'default', or 'unset'."""
|
||||
|
||||
|
||||
class InternalUserSettingsResponse(SettingsResponse):
|
||||
|
|
@ -717,7 +753,7 @@ async def get_sso_settings():
|
|||
Returns a structured object with values and descriptions for UI display.
|
||||
"""
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_config
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -725,59 +761,12 @@ async def get_sso_settings():
|
|||
detail={"error": "Database not connected. Please connect a database."},
|
||||
)
|
||||
|
||||
# Get SSO config from dedicated table
|
||||
# Resolve the effective SSO config: the stored row wins, else the process
|
||||
# environment, else each field's default. Unlike the legacy read path this
|
||||
# does not write os.environ; a GET has no business mutating the environment.
|
||||
sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"})
|
||||
|
||||
# Initialize with defaults
|
||||
sso_settings_dict = {}
|
||||
|
||||
if sso_db_record and sso_db_record.sso_settings:
|
||||
# Load settings from database
|
||||
sso_settings_dict = dict(sso_db_record.sso_settings)
|
||||
|
||||
role_mappings_data = sso_settings_dict.pop("role_mappings", None)
|
||||
role_mappings = None
|
||||
if role_mappings_data:
|
||||
from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings
|
||||
|
||||
if isinstance(role_mappings_data, dict):
|
||||
role_mappings = RoleMappings(**role_mappings_data)
|
||||
elif isinstance(role_mappings_data, RoleMappings):
|
||||
role_mappings = role_mappings_data
|
||||
|
||||
team_mappings_data = sso_settings_dict.pop("team_mappings", None)
|
||||
team_mappings = None
|
||||
if team_mappings_data:
|
||||
from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
|
||||
|
||||
if isinstance(team_mappings_data, dict):
|
||||
team_mappings = TeamMappings(**team_mappings_data)
|
||||
elif isinstance(team_mappings_data, TeamMappings):
|
||||
team_mappings = team_mappings_data
|
||||
|
||||
decrypted_sso_settings_dict = proxy_config._decrypt_and_set_db_env_variables(
|
||||
environment_variables=sso_settings_dict
|
||||
)
|
||||
|
||||
# Build SSO config with database values or environment fallback
|
||||
|
||||
sso_config = SSOConfig(
|
||||
google_client_id=decrypted_sso_settings_dict.get("google_client_id", None),
|
||||
google_client_secret=decrypted_sso_settings_dict.get("google_client_secret", None),
|
||||
microsoft_client_id=decrypted_sso_settings_dict.get("microsoft_client_id", None),
|
||||
microsoft_client_secret=decrypted_sso_settings_dict.get("microsoft_client_secret", None),
|
||||
microsoft_tenant=decrypted_sso_settings_dict.get("microsoft_tenant", None),
|
||||
generic_client_id=decrypted_sso_settings_dict.get("generic_client_id", None),
|
||||
generic_client_secret=decrypted_sso_settings_dict.get("generic_client_secret", None),
|
||||
generic_authorization_endpoint=decrypted_sso_settings_dict.get("generic_authorization_endpoint", None),
|
||||
generic_token_endpoint=decrypted_sso_settings_dict.get("generic_token_endpoint", None),
|
||||
generic_userinfo_endpoint=decrypted_sso_settings_dict.get("generic_userinfo_endpoint", None),
|
||||
proxy_base_url=decrypted_sso_settings_dict.get("proxy_base_url", None),
|
||||
user_email=decrypted_sso_settings_dict.get("user_email"),
|
||||
ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"),
|
||||
role_mappings=role_mappings,
|
||||
team_mappings=team_mappings,
|
||||
)
|
||||
sso_db_settings = dict(sso_db_record.sso_settings) if sso_db_record and sso_db_record.sso_settings else None
|
||||
resolved = resolve_sso_config(sso_db_settings, os.environ)
|
||||
|
||||
# Get the schema for UI display
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -786,11 +775,12 @@ async def get_sso_settings():
|
|||
|
||||
# Convert to dict for response, masking OAuth client secrets so plaintext
|
||||
# is never sent to the UI.
|
||||
sso_dict = mask_sensitive_keys(sso_config.model_dump(), _SSO_SENSITIVE_FIELDS)
|
||||
sso_dict = mask_sensitive_keys(resolved.config.model_dump(), set(SSO_SECRET_FIELDS))
|
||||
|
||||
# Add descriptions to the response
|
||||
result = {
|
||||
"values": sso_dict,
|
||||
"provenance": resolved.provenance,
|
||||
"field_schema": {
|
||||
"description": schema.get("description", ""),
|
||||
"properties": {},
|
||||
|
|
@ -841,21 +831,6 @@ async def update_sso_settings(
|
|||
detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."},
|
||||
)
|
||||
|
||||
# Update environment variables
|
||||
env_var_mapping = {
|
||||
"google_client_id": "GOOGLE_CLIENT_ID",
|
||||
"google_client_secret": "GOOGLE_CLIENT_SECRET",
|
||||
"microsoft_client_id": "MICROSOFT_CLIENT_ID",
|
||||
"microsoft_client_secret": "MICROSOFT_CLIENT_SECRET",
|
||||
"microsoft_tenant": "MICROSOFT_TENANT",
|
||||
"generic_client_id": "GENERIC_CLIENT_ID",
|
||||
"generic_client_secret": "GENERIC_CLIENT_SECRET",
|
||||
"generic_authorization_endpoint": "GENERIC_AUTHORIZATION_ENDPOINT",
|
||||
"generic_token_endpoint": "GENERIC_TOKEN_ENDPOINT",
|
||||
"generic_userinfo_endpoint": "GENERIC_USERINFO_ENDPOINT",
|
||||
"proxy_base_url": "PROXY_BASE_URL",
|
||||
}
|
||||
|
||||
# Read the existing SSO row first so the audit log captures a real
|
||||
# before/after diff. Stored values are encrypted; decrypt them so the
|
||||
# before-snapshot has the same shape as after_value, and rely on
|
||||
|
|
@ -884,8 +859,8 @@ async def update_sso_settings(
|
|||
# Update environment variables in config and in memory
|
||||
sso_data = sso_config.model_dump()
|
||||
for field_name, value in sso_data.items():
|
||||
if field_name in env_var_mapping:
|
||||
env_var_name = env_var_mapping[field_name]
|
||||
if field_name in SSO_FIELD_ENV_VARS:
|
||||
env_var_name = SSO_FIELD_ENV_VARS[field_name]
|
||||
if value:
|
||||
os.environ[env_var_name] = value
|
||||
else:
|
||||
|
|
@ -935,7 +910,7 @@ async def update_sso_settings(
|
|||
else:
|
||||
environment_variables = {}
|
||||
|
||||
env_vars_to_remove = set(env_var_mapping.values())
|
||||
env_vars_to_remove = set(SSO_FIELD_ENV_VARS.values())
|
||||
filtered_env_vars = {
|
||||
key: value for key, value in environment_variables.items() if key not in env_vars_to_remove
|
||||
}
|
||||
|
|
@ -977,12 +952,19 @@ async def get_ui_theme_settings():
|
|||
# Load existing config
|
||||
config = await proxy_config.get_config()
|
||||
|
||||
return await _get_settings_with_schema(
|
||||
result = await _get_settings_with_schema(
|
||||
settings_key="ui_theme_config",
|
||||
settings_class=UIThemeConfig,
|
||||
config=config,
|
||||
)
|
||||
|
||||
stored_values = result.get("values", {})
|
||||
result["values"] = {
|
||||
**stored_values,
|
||||
**{field: _resolve_ui_theme_field(stored_values, field) for field in _UI_THEME_FIELD_ENV_VARS},
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def _validate_public_image_url(value: Optional[str], field_name: str) -> None:
|
||||
"""
|
||||
|
|
@ -1041,13 +1023,6 @@ async def update_ui_theme_settings(
|
|||
config = await proxy_config.get_config()
|
||||
before_theme = config.get("litellm_settings", {}).get("ui_theme_config")
|
||||
|
||||
# Update config with UI theme settings
|
||||
if "general_settings" not in config:
|
||||
config["general_settings"] = {}
|
||||
|
||||
if "environment_variables" not in config:
|
||||
config["environment_variables"] = {}
|
||||
|
||||
# Convert theme config to dict
|
||||
theme_data = theme_config.model_dump(exclude_none=True)
|
||||
|
||||
|
|
@ -1056,55 +1031,29 @@ async def update_ui_theme_settings(
|
|||
config["litellm_settings"] = {}
|
||||
config["litellm_settings"]["ui_theme_config"] = theme_data
|
||||
|
||||
# Update UI_LOGO_PATH environment variable if logo_url is provided
|
||||
# If logo_url is empty string, None, or null, remove the environment variable to use default
|
||||
logo_url = theme_data.get("logo_url")
|
||||
verbose_proxy_logger.debug(f"Updating logo_url: {logo_url}")
|
||||
# UI_LOGO_PATH and LITELLM_FAVICON_URL are the only environment variables
|
||||
# this endpoint owns. A non-empty value sets the var; an empty or missing
|
||||
# one clears it back to the default. Apply to the live process immediately,
|
||||
# then persist only these two keys so an unrelated env var (a YAML/OS value
|
||||
# merged in by get_config) is never snapshotted into the DB.
|
||||
def _clean(url: str | None) -> str | None:
|
||||
return url if url is not None and url.strip() else None
|
||||
|
||||
if (
|
||||
logo_url and isinstance(logo_url, str) and logo_url.strip()
|
||||
): # Check if logo_url exists and is not empty/whitespace
|
||||
config["environment_variables"]["UI_LOGO_PATH"] = logo_url
|
||||
os.environ["UI_LOGO_PATH"] = logo_url
|
||||
verbose_proxy_logger.debug(f"Set UI_LOGO_PATH to: {logo_url}")
|
||||
else:
|
||||
# Remove the environment variable to restore default logo
|
||||
if "UI_LOGO_PATH" in config.get("environment_variables", {}):
|
||||
del config["environment_variables"]["UI_LOGO_PATH"]
|
||||
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from config")
|
||||
if "UI_LOGO_PATH" in os.environ:
|
||||
del os.environ["UI_LOGO_PATH"]
|
||||
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from environment")
|
||||
env_updates: dict[str, str | None] = {
|
||||
"UI_LOGO_PATH": _clean(theme_config.logo_url),
|
||||
"LITELLM_FAVICON_URL": _clean(theme_config.favicon_url),
|
||||
}
|
||||
for env_key, env_value in env_updates.items():
|
||||
if env_value is not None:
|
||||
os.environ[env_key] = env_value
|
||||
else:
|
||||
os.environ.pop(env_key, None)
|
||||
|
||||
# Update LITELLM_FAVICON_URL environment variable if favicon_url is provided
|
||||
favicon_url = theme_data.get("favicon_url")
|
||||
verbose_proxy_logger.debug(f"Updating favicon_url: {favicon_url}")
|
||||
|
||||
if (
|
||||
favicon_url and isinstance(favicon_url, str) and favicon_url.strip()
|
||||
): # Check if favicon_url exists and is not empty/whitespace
|
||||
config["environment_variables"]["LITELLM_FAVICON_URL"] = favicon_url
|
||||
os.environ["LITELLM_FAVICON_URL"] = favicon_url
|
||||
verbose_proxy_logger.debug(f"Set LITELLM_FAVICON_URL to: {favicon_url}")
|
||||
else:
|
||||
# Remove the environment variable to restore default favicon
|
||||
if "LITELLM_FAVICON_URL" in config.get("environment_variables", {}):
|
||||
del config["environment_variables"]["LITELLM_FAVICON_URL"]
|
||||
verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from config")
|
||||
if "LITELLM_FAVICON_URL" in os.environ:
|
||||
del os.environ["LITELLM_FAVICON_URL"]
|
||||
verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from environment")
|
||||
|
||||
# Handle environment variable encryption if needed
|
||||
stored_config = config.copy()
|
||||
if "environment_variables" in stored_config and len(stored_config["environment_variables"]) > 0:
|
||||
# Only encrypt if there are environment variables to encrypt
|
||||
stored_config["environment_variables"] = proxy_config._encrypt_env_variables(
|
||||
environment_variables=stored_config["environment_variables"]
|
||||
)
|
||||
|
||||
# Save the updated config
|
||||
await proxy_config.save_config(new_config=stored_config)
|
||||
# Persist the theme config (litellm_settings). save_config defaults to
|
||||
# include_env_vars=False, so it does not snapshot environment_variables.
|
||||
await proxy_config.save_config(new_config=config)
|
||||
# Persist only the two owned env vars, merged against the existing DB row.
|
||||
await proxy_config.save_environment_variables(env_updates)
|
||||
|
||||
asyncio.create_task(
|
||||
create_config_audit_log(
|
||||
|
|
|
|||
|
|
@ -172,6 +172,7 @@ from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
from prisma.client import TransactionManager
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
|
@ -2922,6 +2923,14 @@ class PrismaClient:
|
|||
return self.db.writer
|
||||
return self.db
|
||||
|
||||
def tx(self) -> "TransactionManager":
|
||||
"""Open an interactive transaction on the writer.
|
||||
|
||||
Callers go through this instead of reaching into ``self.db`` so writer
|
||||
selection and read-replica routing stay encapsulated in the wrapper.
|
||||
"""
|
||||
return cast("TransactionManager", self.db.tx()) # cast-ok: wrappers delegate tx via __getattr__ (untyped)
|
||||
|
||||
def get_request_status(self, payload: Union[dict, SpendLogsPayload]) -> Literal["success", "failure"]:
|
||||
"""
|
||||
Determine if a request was successful or failed based on payload metadata.
|
||||
|
|
@ -6159,6 +6168,9 @@ def create_model_info_response(
|
|||
if model_cost_info is not None:
|
||||
max_input_tokens = coerce_token_limit(model_cost_info.get("max_input_tokens"))
|
||||
max_output_tokens = coerce_token_limit(model_cost_info.get("max_output_tokens"))
|
||||
mode = model_cost_info.get("mode")
|
||||
if isinstance(mode, str):
|
||||
base["mode"] = mode
|
||||
|
||||
if llm_router is not None:
|
||||
configured_input, configured_output = llm_router.get_configured_token_limits(model_id)
|
||||
|
|
|
|||
|
|
@ -4,11 +4,18 @@ Team repository for database operations on LiteLLM_TeamTable.
|
|||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type
|
||||
|
||||
from litellm.models.team import LiteLLM_TeamTable
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.models.team import LiteLLM_TeamTable, Member
|
||||
from litellm.repositories.base_repository import BaseRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
|
||||
_MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member])
|
||||
|
||||
|
||||
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
||||
"""Repository for team database operations."""
|
||||
|
|
@ -46,6 +53,24 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
|||
|
||||
return LiteLLM_TeamTable(**data)
|
||||
|
||||
async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]:
|
||||
"""Return the team's members_with_roles, locking the row FOR UPDATE.
|
||||
|
||||
Must be called inside a transaction so the row lock is held until
|
||||
commit. This serializes concurrent membership writers on the team row
|
||||
so the losing writer appends onto the winner's committed result instead
|
||||
of overwriting it from a stale snapshot.
|
||||
"""
|
||||
rows = await tx.query_raw(
|
||||
'SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = $1 FOR UPDATE',
|
||||
team_id,
|
||||
)
|
||||
raw_value = rows[0]["members_with_roles"] if rows else None
|
||||
parsed = json.loads(raw_value) if isinstance(raw_value, str) else raw_value
|
||||
if not parsed:
|
||||
return []
|
||||
return _MEMBERS_WITH_ROLES_ADAPTER.validate_python(parsed)
|
||||
|
||||
async def find_by_id(self, team_id: str, id_field: str = "team_id") -> Optional[LiteLLM_TeamTable]:
|
||||
return await super().find_by_id(team_id, id_field)
|
||||
|
||||
|
|
|
|||
|
|
@ -494,7 +494,14 @@ async def aresponses(
|
|||
prompt_label=kwargs.get("prompt_label", None),
|
||||
prompt_version=kwargs.get("prompt_version", None),
|
||||
)
|
||||
input = cast(Union[str, ResponseInputParam], merged_input)
|
||||
input = cast(
|
||||
Union[str, ResponseInputParam],
|
||||
ResponsesAPIRequestUtils.merge_prompt_management_input(
|
||||
original_input=input,
|
||||
client_input=client_input,
|
||||
merged_input=merged_input,
|
||||
),
|
||||
)
|
||||
if model != original_model:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
|
||||
kwargs.pop("prompt_id", None)
|
||||
|
|
@ -609,7 +616,14 @@ def _apply_prompt_management_to_responses_call(
|
|||
prompt_label=kwargs.get("prompt_label", None),
|
||||
prompt_version=kwargs.get("prompt_version", None),
|
||||
)
|
||||
input = cast(Union[str, ResponseInputParam], merged_input)
|
||||
input = cast(
|
||||
Union[str, ResponseInputParam],
|
||||
ResponsesAPIRequestUtils.merge_prompt_management_input(
|
||||
original_input=input,
|
||||
client_input=client_input,
|
||||
merged_input=merged_input,
|
||||
),
|
||||
)
|
||||
local_vars["input"] = input
|
||||
local_vars["model"] = model
|
||||
if model != original_model:
|
||||
|
|
|
|||
|
|
@ -19,7 +19,9 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ResponseAPIUsage,
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponseText,
|
||||
|
|
@ -36,6 +38,57 @@ from litellm.types.utils import (
|
|||
class ResponsesAPIRequestUtils:
|
||||
"""Helper utils for constructing ResponseAPI requests"""
|
||||
|
||||
@staticmethod
|
||||
def merge_prompt_management_input(
|
||||
original_input: str | ResponseInputParam,
|
||||
client_input: list[AllMessageValues],
|
||||
merged_input: list[AllMessageValues],
|
||||
) -> list[object]:
|
||||
if isinstance(original_input, str):
|
||||
return [*merged_input]
|
||||
|
||||
original_items = tuple(original_input)
|
||||
client_item_ids = frozenset(id(item) for item in client_input)
|
||||
message_positions = tuple(index for index, item in enumerate(original_items) if id(item) in client_item_ids)
|
||||
|
||||
if len(message_positions) == len(original_items):
|
||||
return [*merged_input]
|
||||
if not message_positions:
|
||||
verbose_logger.warning(
|
||||
"Prompt management hook returned messages without Responses API input messages; merged messages were ignored"
|
||||
)
|
||||
return [*original_items]
|
||||
|
||||
corresponding_messages = len(client_input) == len(merged_input) and all(
|
||||
original.get("role") == merged.get("role")
|
||||
and (not isinstance(original.get("id"), str) or original.get("id") == merged.get("id"))
|
||||
for original, merged in zip(client_input, merged_input)
|
||||
)
|
||||
if corresponding_messages:
|
||||
merged_by_position = dict(zip(message_positions, merged_input))
|
||||
return [
|
||||
merged_by_position[index] if index in merged_by_position else item
|
||||
for index, item in enumerate(original_items)
|
||||
]
|
||||
|
||||
all_messages_preserved = all(any(original is merged for merged in merged_input) for original in client_input)
|
||||
if all_messages_preserved:
|
||||
prefixes = {
|
||||
id(original_items[position]): original_items[
|
||||
message_positions[index - 1] + 1 if index else 0 : position
|
||||
]
|
||||
for index, position in enumerate(message_positions)
|
||||
}
|
||||
trailing_items = original_items[message_positions[-1] + 1 :]
|
||||
return [item for merged in merged_input for item in (*prefixes.get(id(merged), ()), merged)] + list(
|
||||
trailing_items
|
||||
)
|
||||
|
||||
verbose_logger.warning(
|
||||
"Prompt management hook replaced Responses API messages; non-message input items were dropped"
|
||||
)
|
||||
return [*merged_input]
|
||||
|
||||
@staticmethod
|
||||
def _check_valid_arg(
|
||||
supported_params: Optional[List[str]],
|
||||
|
|
|
|||
|
|
@ -725,6 +725,18 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
description="When True, guardrails only receive the latest message for the relevant role (e.g., newest user input pre-call, newest assistant output post-call)",
|
||||
)
|
||||
|
||||
only_scan_new_messages: Optional[bool] = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When True, the guardrail only scans messages that have not already been scanned "
|
||||
"earlier in the same session (identified by litellm_session_id / session_id). "
|
||||
"Message content is hashed per session and cached; only the diff (new or edited "
|
||||
"messages) is sent to the guardrail provider on follow-up calls. Falls back to a "
|
||||
"full scan when the request has no session id or the cache is unavailable. Intended "
|
||||
"for blocking/detection guardrails; not applied when mask_request_content is set."
|
||||
),
|
||||
)
|
||||
|
||||
skip_system_message_in_guardrail: Optional[bool] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -826,6 +838,13 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
"while fail_on_error still governs real Model Armor API errors. Default False blocks them."
|
||||
),
|
||||
)
|
||||
sanitize_error_detail: Optional[bool] = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"For guardrail='model_armor': omit the raw Model Armor response from "
|
||||
"caller-facing errors and logs by default. Set False to restore verbose output."
|
||||
),
|
||||
)
|
||||
|
||||
additional_provider_specific_params: Optional[Dict[str, Any]] = Field(
|
||||
default=None,
|
||||
|
|
|
|||
|
|
@ -173,6 +173,7 @@ class Status1(Enum):
|
|||
cancelled = "cancelled"
|
||||
incomplete = "incomplete"
|
||||
budget_exceeded = "budget_exceeded"
|
||||
queued = "queued"
|
||||
|
||||
|
||||
class InteractionStatusUpdate(BaseModel):
|
||||
|
|
@ -341,6 +342,7 @@ class Status3(Enum):
|
|||
CANCELLED = "cancelled"
|
||||
INCOMPLETE = "incomplete"
|
||||
BUDGET_EXCEEDED = "budget_exceeded"
|
||||
QUEUED = "queued"
|
||||
|
||||
|
||||
class ModelOption(RootModel[str]):
|
||||
|
|
|
|||
|
|
@ -20,6 +20,13 @@ class ModelArmorGuardrailConfigModel(GuardrailConfigModel):
|
|||
default=True,
|
||||
description="Whether to fail the request if Model Armor encounters an error",
|
||||
)
|
||||
sanitize_error_detail: Optional[bool] = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Omit the raw Model Armor response from caller-facing errors and logs "
|
||||
"by default. Set False to restore verbose output."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
|
|
|
|||
|
|
@ -148,6 +148,10 @@ class SSOConfig(LiteLLMPydanticObjectBase):
|
|||
default=None,
|
||||
description="User info endpoint URL for generic OAuth provider",
|
||||
)
|
||||
generic_scope: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Space-separated OAuth scopes requested from the generic provider, e.g. 'openid email profile'",
|
||||
)
|
||||
|
||||
# Common settings
|
||||
proxy_base_url: Optional[str] = Field(
|
||||
|
|
|
|||
|
|
@ -10,12 +10,16 @@ class ModelInfoMetadata(TypedDict):
|
|||
|
||||
|
||||
class ModelInfoResponse(TypedDict):
|
||||
"""OpenAI-compatible model object. `metadata` is present only when the
|
||||
endpoint is called with include_metadata=true.
|
||||
"""OpenAI-compatible model object. `mode`, `max_input_tokens`, and
|
||||
`max_output_tokens` are attached when the cost map knows them; `metadata`
|
||||
is present only when the endpoint is called with include_metadata=true.
|
||||
"""
|
||||
|
||||
id: str
|
||||
object: Literal["model"]
|
||||
created: int
|
||||
owned_by: str
|
||||
mode: NotRequired[str]
|
||||
max_input_tokens: NotRequired[int]
|
||||
max_output_tokens: NotRequired[int]
|
||||
metadata: NotRequired[ModelInfoMetadata]
|
||||
|
|
|
|||
|
|
@ -1864,7 +1864,9 @@ def client(original_function):
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
setattr(e, "num_retries", num_retries) ## IMPORTANT: returns the deployment's num_retries to the router
|
||||
deployment_num_retries = kwargs.get("num_retries")
|
||||
if deployment_num_retries is not None:
|
||||
setattr(e, "num_retries", deployment_num_retries)
|
||||
|
||||
timeout = _get_wrapper_timeout(kwargs=kwargs, exception=e)
|
||||
setattr(e, "timeout", timeout)
|
||||
|
|
|
|||
|
|
@ -17644,6 +17644,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
|
|
@ -18311,6 +18366,60 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
|
|
@ -19663,6 +19772,63 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"rpm": 15,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 250000,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3-flash-preview": {
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -19769,6 +19935,63 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/pricing/gemini-3",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-omni-flash-preview": {
|
||||
"input_cost_per_audio_token": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
|
|
@ -20049,6 +20272,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini-3.6-flash": {
|
||||
"cache_read_input_token_cost": 1.5e-07,
|
||||
"cache_read_input_token_cost_flex": 7.5e-08,
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"input_cost_per_token_batches": 7.5e-07,
|
||||
"input_cost_per_token_flex": 7.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 7.5e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"output_cost_per_token_batches": 3.75e-06,
|
||||
"output_cost_per_token_flex": 3.75e-06,
|
||||
"source": "https://ai.google.dev/pricing/gemini-3",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"input_cost_per_token_priority": 2.7e-06,
|
||||
"output_cost_per_token_priority": 1.35e-05,
|
||||
"cache_read_input_token_cost_priority": 2.7e-07,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"gemini/gemini-2.5-pro-preview-tts": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
|
|
@ -37323,6 +37601,61 @@
|
|||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/gemini-3.5-flash-lite": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"cache_read_input_token_cost_flex": 2e-08,
|
||||
"cache_read_input_token_cost_priority": 5e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 1.25e-06,
|
||||
"output_cost_per_token_flex": 1.25e-06,
|
||||
"output_cost_per_token_priority": 4.5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_native_streaming": true,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_low": 0.014,
|
||||
"search_context_size_medium": 0.014,
|
||||
"search_context_size_high": 0.014
|
||||
},
|
||||
"web_search_billing_unit": "per_query"
|
||||
},
|
||||
"vertex_ai/deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm"
|
||||
version = "1.94.0"
|
||||
version = "1.95.0"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10, <3.15"
|
||||
|
|
@ -62,7 +62,7 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.28.1,<2.0",
|
||||
"litellm-proxy-extras==0.4.79",
|
||||
"litellm-proxy-extras==0.4.80",
|
||||
"litellm-enterprise==0.1.51",
|
||||
"RestrictedPython>=8.1,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
|
|
@ -117,6 +117,14 @@ stt-nvidia-riva = [
|
|||
"numpy>=1.26.0",
|
||||
]
|
||||
google = ["google-cloud-aiplatform>=1.133.0,<2.0"]
|
||||
bedrock-realtime = [
|
||||
# Bedrock Nova Sonic realtime (speech-to-speech) uses the
|
||||
# InvokeModelWithBidirectionalStream API, which boto3 cannot do. This
|
||||
# experimental AWS SDK (with its smithy-* deps, pulled transitively)
|
||||
# provides the bidirectional stream; imported lazily in the realtime
|
||||
# handler so litellm core stays usable without it.
|
||||
"aws-sdk-bedrock-runtime>=0.7.0,<0.8.0; python_version >= '3.12'",
|
||||
]
|
||||
proxy-runtime = [
|
||||
# Historically bundled in the proxy Docker images via requirements.txt.
|
||||
# Keep these in a dedicated extra so uv-based images preserve the same
|
||||
|
|
@ -289,7 +297,7 @@ members = ["enterprise", "litellm-proxy-extras"]
|
|||
profile = "black"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.94.0"
|
||||
version = "1.95.0"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
{
|
||||
"include": ["litellm"],
|
||||
"ignore": [],
|
||||
"exclude": ["**/node_modules", "**/__pycache__", "tests/e2e/claude_code", "litellm/types/utils.py", "litellm/proxy/_types.py"],
|
||||
"exclude": ["**/node_modules", "**/__pycache__", "tests/e2e/claude_code", "tests/e2e/ui", "litellm/types/utils.py", "litellm/proxy/_types.py"],
|
||||
"pythonVersion": "3.12",
|
||||
"typeCheckingMode": "strict",
|
||||
"enableTypeIgnoreComments": false,
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@
|
|||
"limit": 33
|
||||
},
|
||||
"DTZ005": {
|
||||
"limit": 244
|
||||
"limit": 241
|
||||
},
|
||||
"DTZ006": {
|
||||
"limit": 13
|
||||
|
|
|
|||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@
|
|||
# gating CI checks, so a clean run means a green CI lint:
|
||||
# - litellm/ Python staged -> `make lint` (test-linting.yml's lint job)
|
||||
# - tests/e2e Python staged -> `make lint-e2e-basedpyright` (test-linting.yml's e2e type-check step)
|
||||
# + raw HTTP client ban (test-code-quality.yml's check_e2e_no_raw_requests)
|
||||
# - dashboard staged -> prettier + eslint + lint budgets (test-litellm-ui-build.yml's frontend-lint)
|
||||
# - proxy/types staged -> regenerate dashboard API types and fail on drift (check-ui-api-types.yml)
|
||||
#
|
||||
|
|
@ -112,6 +113,12 @@ if [ -n "$e2e_py_files" ] && [ -z "$litellm_py_files" ]; then
|
|||
make lint-e2e-basedpyright || { echo "✗ tests/e2e basedpyright failed. Fix the errors above, then re-run make pre-commit." >&2; status=1; }
|
||||
fi
|
||||
|
||||
if [ -n "$e2e_py_files" ]; then
|
||||
echo "pre-commit: checking tests/e2e raw HTTP client ban (check_e2e_no_raw_requests)"
|
||||
uv run --no-sync python tests/code_coverage_tests/check_e2e_no_raw_requests.py \
|
||||
|| { echo "✗ Raw HTTP client import in tests/e2e. Route the call through tests/e2e/e2e_http.py, then re-run make pre-commit." >&2; status=1; }
|
||||
fi
|
||||
|
||||
if [ -n "$ui_prettier_files" ] || [ -n "$ui_eslint_files" ]; then
|
||||
echo "pre-commit: linting dashboard (prettier + eslint + lint budgets)"
|
||||
if [ ! -d ui/litellm-dashboard/node_modules ]; then
|
||||
|
|
|
|||
81
tests/code_coverage_tests/check_e2e_no_raw_requests.py
Normal file
81
tests/code_coverage_tests/check_e2e_no_raw_requests.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
"""tests/e2e routes every HTTP call through the typed transport (e2e_http.py), so
|
||||
raw HTTP client imports (requests, urllib.request, httpx, aiohttp, http.client) are
|
||||
banned in suite code. Importing requests' exception types for catching is fine
|
||||
anywhere; a small allowlist grandfathers the files that legitimately make raw calls
|
||||
(the transport itself, the root conftest liveness probe, and the claude_code version
|
||||
resolver's constant registry URL fetch). Referenced by tests/e2e/CLAUDE.md."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
E2E_DIR = Path(__file__).resolve().parents[1] / "e2e"
|
||||
|
||||
BANNED_MODULES = ("requests", "urllib.request", "http.client", "httpx", "aiohttp")
|
||||
|
||||
ALLOWED_RAW_CLIENT_FILES = {
|
||||
"e2e_http.py": ("requests",),
|
||||
"conftest.py": ("requests",),
|
||||
"claude_code/pr_gate_version_resolver.py": ("urllib.request",),
|
||||
}
|
||||
|
||||
EXCEPTION_ONLY_NAMES = frozenset({"RequestException", "ConnectionError", "Timeout", "HTTPError"})
|
||||
|
||||
|
||||
def _is_banned(module: str) -> bool:
|
||||
return any(module == banned or module.startswith(banned + ".") for banned in BANNED_MODULES)
|
||||
|
||||
|
||||
def _banned_imports(tree: ast.Module) -> tuple[tuple[str, int], ...]:
|
||||
plain = tuple(
|
||||
(alias.name, node.lineno)
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.Import)
|
||||
for alias in node.names
|
||||
if _is_banned(alias.name)
|
||||
)
|
||||
from_imports = tuple(
|
||||
(node.module, node.lineno)
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.ImportFrom)
|
||||
and node.module is not None
|
||||
and _is_banned(node.module)
|
||||
and not all(alias.name in EXCEPTION_ONLY_NAMES for alias in node.names)
|
||||
)
|
||||
return plain + from_imports
|
||||
|
||||
|
||||
def _violations_in(path: Path) -> tuple[str, ...]:
|
||||
relative = path.relative_to(E2E_DIR).as_posix()
|
||||
allowed = ALLOWED_RAW_CLIENT_FILES.get(relative, ())
|
||||
tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
|
||||
return tuple(
|
||||
f"tests/e2e/{relative}:{lineno}: raw HTTP client import '{module}'"
|
||||
for module, lineno in _banned_imports(tree)
|
||||
if module not in allowed
|
||||
)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
violations = tuple(
|
||||
violation
|
||||
for path in sorted(E2E_DIR.rglob("*.py"))
|
||||
for violation in _violations_in(path)
|
||||
)
|
||||
for violation in violations:
|
||||
print(violation)
|
||||
if violations:
|
||||
print(
|
||||
f"\n{len(violations)} raw HTTP client import(s) in tests/e2e. "
|
||||
"Route the call through tests/e2e/e2e_http.py (get_external for absolute "
|
||||
"third-party URLs) so it gets the typed Result handling."
|
||||
)
|
||||
return 1
|
||||
print("tests/e2e raw HTTP client check passed")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
|
|
@ -55,6 +55,8 @@ IGNORE_FUNCTIONS = [
|
|||
"_freeze_for_dedupe", # OTEL: max depth set (default 16, _FREEZE_MAX_DEPTH); fails closed by returning repr(value) at the cap.
|
||||
"apply_json_merge_patch", # max depth set (_MAX_MERGE_DEPTH=64); fails closed by raising ValueError at the cap.
|
||||
"_filter_mcp_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the MCP call at the cap.
|
||||
"_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap.
|
||||
"_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -13,13 +13,16 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
|
|||
- `realtime/` - realtime websocket sessions, including the pipecat audio path
|
||||
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
|
||||
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials (API surface; not Playwright)
|
||||
- `a2a/` - the A2A (agent-to-agent) surface: admin registration via `/v1/agents`, proxy-fronted card discovery at `/.well-known/agent-card.json`, and JSON-RPC `message/send` invocation, driving agents backed by the litellm completion bridge (a real provider) and asserting protocol-version normalization (0.3 vs 1.0)
|
||||
- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server only (see "MCP suite: real Datadog only" below)
|
||||
- `logging/` - logging-integration delivery (datadog and friends)
|
||||
- `security/` - secret handling and log-leak protection
|
||||
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
|
||||
- `load/` - throughput/performance under concurrency: drives real concurrent traffic through the whole stack with Locust and asserts a throughput SLO; marked `load` so the parent conftest collects it last and it never perturbs latency-sensitive suites
|
||||
- `load/` - throughput/performance under concurrency: drives real concurrent traffic through the whole stack with Locust and asserts a throughput SLO; marked `load` so the parent conftest collects it last and it never perturbs latency-sensitive suites. Also home of the weekly session-anomaly test (`test_weekly_session_anomaly_e2e.py`): Claude Code-shaped multi-turn sessions against real providers with ceilings on error rate, cache read/write, turn time, and spend; additionally marked `weekly` and deselected unless `E2E_WEEKLY_ANOMALY` is set, because it spends real provider money (driven by `.github/workflows/weekly_load_anomaly.yml`)
|
||||
- `other/` - the holding-pen suite for the `other.*` registry cluster with no home of its own yet: the master-key auth gate and the process-lifecycle health probes (liveness, public readiness, authenticated readiness diagnostics). Promote a cluster out once it is large/stable enough for its own suite
|
||||
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
|
||||
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees. The HTTP probes ride the shared transport (`ProxyClient.count_tokens` / `ProxyClient.messages`); the CLI-driving path stays bespoke
|
||||
- `ui/` - the Admin UI browser suite: Playwright in TypeScript, driving the dashboard served by a live proxy on port 4000 (seeded postgres + mock LLM upstream; see its `run_e2e.sh`). It is a self-contained npm package with its own lockfile and does not use the Python harness, pytest markers, or the shared transport; the Python rules in this file (typed models, `Result` unions, basedpyright zero-error gate) do not apply inside it. Its only Python file, `fixtures/mock_llm_server/server.py`, is excluded from the e2e basedpyright gate via the root `pyrightconfig.json`
|
||||
|
||||
## MCP suite: real Datadog only
|
||||
|
||||
|
|
@ -130,7 +133,7 @@ reliability.<behavior>.<variant>.<assertion>
|
|||
behavior : fallback | retry | cooldown | timeout | routing | cache | circuit_breaker | perf
|
||||
variant : <trigger> 5xx | context_window | content_policy | 429 | timeout
|
||||
<strategy> simple_shuffle | usage_based | latency_based | cost_based | least_busy
|
||||
<dimension> latency | throughput (perf only; SLO/threshold assertion, not binary)
|
||||
<dimension> latency | throughput | session_anomaly (perf only; SLO/threshold assertion, not binary)
|
||||
assertion : routes_to_fallback | succeeds_within_retries | picks_under_tpm | returns_cached
|
||||
| trips_then_recovers | under_slo
|
||||
e.g. reliability.fallback.context_window.routes_to_fallback exercised_on=[chat_completions]
|
||||
|
|
|
|||
292
tests/e2e/a2a/a2a_client.py
Normal file
292
tests/e2e/a2a/a2a_client.py
Normal file
|
|
@ -0,0 +1,292 @@
|
|||
"""Client for the proxy's A2A (agent-to-agent) surface.
|
||||
|
||||
An A2A agent is registered admin-side via POST /v1/agents with an agent card and
|
||||
litellm_params; the proxy fronts it at /a2a/{id}, serving a proxy-owned agent card
|
||||
at /.well-known/agent-card.json and accepting A2A JSON-RPC calls at /a2a/{id}. This
|
||||
suite registers agents backed by the litellm_completion_bridge (custom_llm_provider
|
||||
+ model), so message/send runs a real provider completion and comes back in the
|
||||
agent's pinned A2A protocol version. The A2A request/response models are co-located
|
||||
here because only this suite uses them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from e2e_http import NoBody, Result, get_external, is_ok
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
|
||||
class A2ACapabilities(BaseModel):
|
||||
streaming: bool | None = None
|
||||
push_notifications: bool | None = Field(default=None, serialization_alias="pushNotifications")
|
||||
|
||||
|
||||
class A2ASkill(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
description: str
|
||||
tags: list[str]
|
||||
examples: list[str] | None = None
|
||||
|
||||
|
||||
class A2AProvider(BaseModel):
|
||||
organization: str
|
||||
url: str
|
||||
|
||||
|
||||
class AgentCardParams(BaseModel):
|
||||
"""The upstream agent card an admin registers. `protocolVersion` is the field the
|
||||
proxy validates against SUPPORTED_A2A_PROTOCOL_VERSIONS on registration."""
|
||||
|
||||
protocol_version: str = Field(serialization_alias="protocolVersion")
|
||||
name: str
|
||||
description: str
|
||||
version: str
|
||||
url: str | None = None
|
||||
capabilities: A2ACapabilities = A2ACapabilities()
|
||||
skills: list[A2ASkill]
|
||||
default_input_modes: list[str] = Field(default=["text"], serialization_alias="defaultInputModes")
|
||||
default_output_modes: list[str] = Field(default=["text"], serialization_alias="defaultOutputModes")
|
||||
preferred_transport: str | None = Field(default=None, serialization_alias="preferredTransport")
|
||||
|
||||
|
||||
class UpstreamAgentCard(BaseModel):
|
||||
"""A real published agent card parsed from a public /.well-known endpoint. Keys on
|
||||
the A2A wire aliases so `model_validate_json` reads the served JSON and
|
||||
`model_dump(by_alias=True)` re-emits it unchanged for verbatim registration; it is
|
||||
only ever fetched-and-validated, never hand-constructed, so aliasing on the wire
|
||||
names does not affect any call site."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
protocol_version: str = Field(alias="protocolVersion")
|
||||
name: str
|
||||
description: str
|
||||
version: str
|
||||
url: str
|
||||
provider: A2AProvider | None = None
|
||||
documentation_url: str | None = Field(default=None, alias="documentationUrl")
|
||||
capabilities: A2ACapabilities = A2ACapabilities()
|
||||
skills: list[A2ASkill]
|
||||
default_input_modes: list[str] = Field(default=["text"], alias="defaultInputModes")
|
||||
default_output_modes: list[str] = Field(default=["text"], alias="defaultOutputModes")
|
||||
preferred_transport: str | None = Field(default=None, alias="preferredTransport")
|
||||
|
||||
|
||||
class A2ABridgeParams(BaseModel):
|
||||
"""litellm_params that route the agent through the completion bridge: an A2A
|
||||
message/send is transformed into a litellm.acompletion against this provider."""
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
custom_llm_provider: str
|
||||
model: str
|
||||
|
||||
|
||||
class AgentRegisterBody(BaseModel):
|
||||
agent_name: str
|
||||
agent_card_params: AgentCardParams | UpstreamAgentCard
|
||||
litellm_params: A2ABridgeParams | None = None
|
||||
|
||||
|
||||
class A2ASecurityScheme(BaseModel):
|
||||
type: str
|
||||
scheme: str
|
||||
|
||||
|
||||
class A2AInterface(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
url: str
|
||||
protocol_version: str | None = Field(default=None, alias="protocolVersion")
|
||||
|
||||
|
||||
class ServedAgentCard(BaseModel):
|
||||
"""The proxy-owned card, either nested under a registration response's
|
||||
`agent_card_params` or served raw at /.well-known/agent-card.json. The proxy
|
||||
rewrites `url`/`supportedInterfaces` to itself and replaces the security scheme
|
||||
with its own virtual-key bearer scheme."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
protocol_version: str = Field(alias="protocolVersion")
|
||||
name: str
|
||||
url: str | None = None
|
||||
security_schemes: dict[str, A2ASecurityScheme] | None = Field(default=None, alias="securitySchemes")
|
||||
security: list[dict[str, list[str]]] | None = None
|
||||
supported_interfaces: list[A2AInterface] | None = Field(default=None, alias="supportedInterfaces")
|
||||
|
||||
|
||||
class AgentResponse(BaseModel):
|
||||
agent_id: str
|
||||
agent_name: str
|
||||
agent_card_params: ServedAgentCard
|
||||
|
||||
|
||||
class A2ATextPart(BaseModel):
|
||||
kind: str = "text"
|
||||
text: str
|
||||
|
||||
|
||||
class A2ASearchPropertiesParams(BaseModel):
|
||||
"""The strict param schema of the published property agent's `search_properties`
|
||||
skill (unknown keys are rejected upstream), so a natural-language query like
|
||||
"properties for sale in SF under $2M" is expressed as typed fields."""
|
||||
|
||||
un_locode: str | None = None
|
||||
service_type: str | None = None
|
||||
property_type: str | None = None
|
||||
bedrooms_min: int | None = None
|
||||
asking_price_max: float | None = None
|
||||
limit: int | None = None
|
||||
|
||||
|
||||
class A2ASkillInvocation(BaseModel):
|
||||
skill: str
|
||||
params: A2ASearchPropertiesParams
|
||||
|
||||
|
||||
class A2ADataPart(BaseModel):
|
||||
kind: str = "data"
|
||||
data: A2ASkillInvocation
|
||||
|
||||
|
||||
class A2AOutboundMessage(BaseModel):
|
||||
role: str = "user"
|
||||
parts: list[A2ATextPart | A2ADataPart]
|
||||
message_id: str = Field(serialization_alias="messageId")
|
||||
|
||||
|
||||
class A2AMessageSendParams(BaseModel):
|
||||
message: A2AOutboundMessage
|
||||
|
||||
|
||||
class A2AJsonRpcRequest(BaseModel):
|
||||
jsonrpc: str = "2.0"
|
||||
id: str
|
||||
method: str = "message/send"
|
||||
params: A2AMessageSendParams
|
||||
|
||||
|
||||
class A2AResponsePart(BaseModel):
|
||||
kind: str | None = None
|
||||
text: str | None = None
|
||||
|
||||
|
||||
class A2AResponseMessage(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
message_id: str | None = Field(default=None, alias="messageId")
|
||||
role: str | None = None
|
||||
parts: list[A2AResponsePart] = []
|
||||
|
||||
|
||||
class A2ATaskStatus(BaseModel):
|
||||
state: str | None = None
|
||||
message: A2AResponseMessage | None = None
|
||||
|
||||
|
||||
class A2AResult(BaseModel):
|
||||
"""A message/send result. In 0.3 the message fields sit directly on the result
|
||||
(`kind`/`role`/`parts`); in 1.0 they are nested under `message`; a real agent that
|
||||
runs a task replies with a `task` whose agent text lives on `status.message`.
|
||||
`text` reads the agent's reply from whichever shape the served version produced."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
kind: str | None = None
|
||||
role: str | None = None
|
||||
message_id: str | None = Field(default=None, alias="messageId")
|
||||
parts: list[A2AResponsePart] = []
|
||||
message: A2AResponseMessage | None = None
|
||||
status: A2ATaskStatus | None = None
|
||||
|
||||
@property
|
||||
def text(self) -> str:
|
||||
if self.message is not None:
|
||||
parts = self.message.parts
|
||||
elif self.parts:
|
||||
parts = self.parts
|
||||
elif self.status is not None and self.status.message is not None:
|
||||
parts = self.status.message.parts
|
||||
else:
|
||||
parts = []
|
||||
return "".join(part.text or "" for part in parts)
|
||||
|
||||
@property
|
||||
def is_nested_v1_shape(self) -> bool:
|
||||
return self.message is not None
|
||||
|
||||
|
||||
class A2AError(BaseModel):
|
||||
code: int
|
||||
message: str
|
||||
|
||||
|
||||
class A2AResponse(BaseModel):
|
||||
jsonrpc: str
|
||||
id: str | None = None
|
||||
result: A2AResult | None = None
|
||||
error: A2AError | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class A2AClient:
|
||||
proxy: ProxyClient
|
||||
|
||||
def register_agent(self, body: AgentRegisterBody) -> Result[AgentResponse]:
|
||||
return self.proxy.transport.post(
|
||||
"/v1/agents",
|
||||
headers=self.proxy.transport.master,
|
||||
json=body,
|
||||
response_type=AgentResponse,
|
||||
)
|
||||
|
||||
def get_agent(self, agent_id: str) -> Result[AgentResponse]:
|
||||
return self.proxy.transport.get(
|
||||
f"/v1/agents/{agent_id}",
|
||||
headers=self.proxy.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=AgentResponse,
|
||||
)
|
||||
|
||||
def delete_agent(self, agent_id: str) -> None:
|
||||
result = self.proxy.transport.delete(
|
||||
f"/v1/agents/{agent_id}",
|
||||
headers=self.proxy.transport.master,
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
if not is_ok(result):
|
||||
warnings.warn(f"delete_agent({agent_id!r}) failed: {result}", stacklevel=2)
|
||||
|
||||
def agent_card(self, agent_id: str, key: str) -> Result[ServedAgentCard]:
|
||||
return self.proxy.transport.get(
|
||||
f"/a2a/{agent_id}/.well-known/agent-card.json",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
params=NoBody(),
|
||||
response_type=ServedAgentCard,
|
||||
)
|
||||
|
||||
def send_message(self, agent_id: str, key: str, body: A2AJsonRpcRequest) -> Result[A2AResponse]:
|
||||
return self.proxy.transport.post(
|
||||
f"/a2a/{agent_id}",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=body,
|
||||
response_type=A2AResponse,
|
||||
)
|
||||
|
||||
|
||||
def build_a2a_client(proxy: ProxyClient) -> A2AClient:
|
||||
return A2AClient(proxy=proxy)
|
||||
|
||||
|
||||
def fetch_agent_card(url: str, *, timeout: float = 20.0) -> Result[UpstreamAgentCard]:
|
||||
"""Fetch a live A2A agent card from its /.well-known endpoint and parse it into the
|
||||
registration model, so a test can register a real published card verbatim rather
|
||||
than a hand-rolled one."""
|
||||
return get_external(url, response_type=UpstreamAgentCard, timeout=timeout)
|
||||
17
tests/e2e/a2a/conftest.py
Normal file
17
tests/e2e/a2a/conftest.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
"""A2A suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. A2AClient holds the shared ProxyClient,
|
||||
so the `resources` fixture cleans up keys this suite creates; agents are torn down
|
||||
via `resources.defer(...)` in each test.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from a2a_client import A2AClient, build_a2a_client
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client(proxy: ProxyClient) -> A2AClient:
|
||||
return build_a2a_client(proxy)
|
||||
202
tests/e2e/a2a/test_a2a_agent_e2e.py
Normal file
202
tests/e2e/a2a/test_a2a_agent_e2e.py
Normal file
|
|
@ -0,0 +1,202 @@
|
|||
"""A2A agents end to end, against a live proxy.
|
||||
|
||||
An admin registers an agent whose card pins an A2A protocol version and whose
|
||||
litellm_params route it through the completion bridge; a caller then discovers the
|
||||
proxy-owned card and drives it over A2A JSON-RPC. These tests assert the recorded
|
||||
state (the agent persists, a spend row lands) and the enforced behavior (the served
|
||||
card points back at the proxy, message/send returns a real completion in the pinned
|
||||
protocol version, and an unsupported version is refused at registration).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from a2a_client import (
|
||||
A2ABridgeParams,
|
||||
A2AClient,
|
||||
A2ADataPart,
|
||||
A2AJsonRpcRequest,
|
||||
A2AMessageSendParams,
|
||||
A2AOutboundMessage,
|
||||
A2ASearchPropertiesParams,
|
||||
A2ASkill,
|
||||
A2ASkillInvocation,
|
||||
A2ATextPart,
|
||||
AgentCardParams,
|
||||
AgentRegisterBody,
|
||||
AgentResponse,
|
||||
fetch_agent_card,
|
||||
)
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import Result, UnknownApiError, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
BRIDGE = A2ABridgeParams(custom_llm_provider="anthropic", model="claude-haiku-4-5")
|
||||
|
||||
MOVEHOME_AGENT_CARD_URL = "https://movehome.org/.well-known/agent.json"
|
||||
MOVEHOME_ORIGIN = "https://movehome.org"
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _register(client: A2AClient, resources: ResourceManager, protocol_version: str) -> AgentResponse:
|
||||
marker = unique_marker()
|
||||
body = AgentRegisterBody(
|
||||
agent_name=f"e2e-a2a-{marker}",
|
||||
agent_card_params=AgentCardParams(
|
||||
protocol_version=protocol_version,
|
||||
name=f"E2E A2A {marker}",
|
||||
description="e2e agent backed by the litellm completion bridge",
|
||||
version="1.0.0",
|
||||
skills=[A2ASkill(id="chat", name="Chat", description="general chat", tags=["chat"])],
|
||||
),
|
||||
litellm_params=BRIDGE,
|
||||
)
|
||||
agent = unwrap(client.register_agent(body))
|
||||
resources.defer(lambda: client.delete_agent(agent.agent_id))
|
||||
return agent
|
||||
|
||||
|
||||
def _register_rejection(client: A2AClient, protocol_version: str) -> Result[AgentResponse]:
|
||||
marker = unique_marker()
|
||||
body = AgentRegisterBody(
|
||||
agent_name=f"e2e-a2a-bad-{marker}",
|
||||
agent_card_params=AgentCardParams(
|
||||
protocol_version=protocol_version,
|
||||
name=f"E2E A2A bad {marker}",
|
||||
description="rejected at registration",
|
||||
version="1.0.0",
|
||||
skills=[A2ASkill(id="chat", name="Chat", description="c", tags=["chat"])],
|
||||
),
|
||||
litellm_params=BRIDGE,
|
||||
)
|
||||
return client.register_agent(body)
|
||||
|
||||
|
||||
def _ask(text: str) -> A2AJsonRpcRequest:
|
||||
return A2AJsonRpcRequest(
|
||||
id=f"e2e-{unique_marker()}",
|
||||
params=A2AMessageSendParams(
|
||||
message=A2AOutboundMessage(parts=[A2ATextPart(text=text)], message_id=unique_marker())
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class TestA2AAgentLifecycle:
|
||||
@pytest.mark.covers("other.a2a.register.persists")
|
||||
def test_register_persists(self, client: A2AClient, resources: ResourceManager) -> None:
|
||||
agent = _register(client, resources, "0.3")
|
||||
fetched = unwrap(client.get_agent(agent.agent_id))
|
||||
assert fetched.agent_id == agent.agent_id
|
||||
assert fetched.agent_name == agent.agent_name
|
||||
assert fetched.agent_card_params.protocol_version == "0.3"
|
||||
|
||||
@pytest.mark.covers("other.a2a.register.semver_version_accepted")
|
||||
def test_semver_protocol_version_registers_and_serves(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
agent = _register(client, resources, "0.3.0")
|
||||
assert agent.agent_card_params.protocol_version == "0.3"
|
||||
card = unwrap(client.agent_card(agent.agent_id, scoped_key))
|
||||
assert card.protocol_version == "0.3"
|
||||
assert card.supported_interfaces is not None
|
||||
assert card.supported_interfaces[0].protocol_version == "0.3"
|
||||
result = unwrap(client.send_message(agent.agent_id, scoped_key, _ask("Say hi in one word"))).result
|
||||
assert result is not None
|
||||
assert result.text != ""
|
||||
|
||||
@pytest.mark.covers("other.a2a.message_send.real_world_agent_replies")
|
||||
def test_real_world_agent_replies_to_property_query(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
upstream = unwrap(fetch_agent_card(MOVEHOME_AGENT_CARD_URL)).model_copy(update={"url": MOVEHOME_ORIGIN})
|
||||
assert upstream.protocol_version == "0.3.0"
|
||||
marker = unique_marker()
|
||||
body = AgentRegisterBody(agent_name=f"e2e-a2a-real-{marker}", agent_card_params=upstream)
|
||||
agent = unwrap(client.register_agent(body))
|
||||
resources.defer(lambda: client.delete_agent(agent.agent_id))
|
||||
assert agent.agent_card_params.protocol_version == "0.3"
|
||||
request = A2AJsonRpcRequest(
|
||||
id=f"e2e-{unique_marker()}",
|
||||
params=A2AMessageSendParams(
|
||||
message=A2AOutboundMessage(
|
||||
parts=[
|
||||
A2ADataPart(
|
||||
data=A2ASkillInvocation(
|
||||
skill="search_properties",
|
||||
params=A2ASearchPropertiesParams(un_locode="USSFO", service_type="sale", asking_price_max=2_000_000, limit=3),
|
||||
)
|
||||
)
|
||||
],
|
||||
message_id=unique_marker(),
|
||||
)
|
||||
),
|
||||
)
|
||||
response = unwrap(client.send_message(agent.agent_id, scoped_key, request))
|
||||
assert response.error is None
|
||||
assert response.result is not None
|
||||
assert response.result.text.strip() != ""
|
||||
|
||||
@pytest.mark.covers("other.a2a.discovery.proxy_fronted_card")
|
||||
def test_discovery_card_is_proxy_fronted(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
agent = _register(client, resources, "0.3")
|
||||
card = unwrap(client.agent_card(agent.agent_id, scoped_key))
|
||||
assert card.url is not None and card.url.endswith(f"/a2a/{agent.agent_id}")
|
||||
assert card.security_schemes is not None
|
||||
scheme = next(iter(card.security_schemes.values()))
|
||||
assert scheme.scheme == "bearer"
|
||||
assert card.supported_interfaces is not None
|
||||
assert card.supported_interfaces[0].url == card.url
|
||||
|
||||
@pytest.mark.covers("other.a2a.message_send.bridge_invokes")
|
||||
def test_message_send_runs_completion_bridge(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
agent = _register(client, resources, "0.3")
|
||||
request = _ask("Reply with exactly the word PONG and nothing else")
|
||||
response = unwrap(client.send_message(agent.agent_id, scoped_key, request))
|
||||
assert response.error is None
|
||||
assert response.result is not None
|
||||
assert "PONG" in response.result.text.upper()
|
||||
|
||||
rows = client.proxy.poll_logs_for_request_id(request.id)
|
||||
assert rows, f"no spend log row landed for a2a request {request.id}"
|
||||
assert rows[0].call_type == "asend_message"
|
||||
assert rows[0].model == f"a2a_agent/{agent.agent_card_params.name}"
|
||||
|
||||
@pytest.mark.covers("other.a2a.version.serves_pinned_0_3")
|
||||
def test_pinned_v0_3_serves_flat_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
agent = _register(client, resources, "0.3")
|
||||
request = _ask("Say hi in one word")
|
||||
result = unwrap(client.send_message(agent.agent_id, scoped_key, request)).result
|
||||
assert result is not None
|
||||
assert not result.is_nested_v1_shape
|
||||
assert result.kind == "message"
|
||||
assert result.role == "agent"
|
||||
assert result.text != ""
|
||||
|
||||
@pytest.mark.covers("other.a2a.version.serves_pinned_1_0")
|
||||
def test_pinned_v1_0_serves_nested_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None:
|
||||
agent = _register(client, resources, "1.0")
|
||||
request = _ask("Say hi in one word")
|
||||
result = unwrap(client.send_message(agent.agent_id, scoped_key, request)).result
|
||||
assert result is not None
|
||||
assert result.is_nested_v1_shape
|
||||
assert result.message is not None
|
||||
assert result.message.role == "ROLE_AGENT"
|
||||
assert result.text != ""
|
||||
|
||||
@pytest.mark.covers("other.a2a.register.unsupported_version_rejected")
|
||||
def test_unsupported_protocol_version_rejected(self, client: A2AClient) -> None:
|
||||
result = _register_rejection(client, "9.9")
|
||||
match result:
|
||||
case UnknownApiError(status_code=status, body=detail):
|
||||
assert status == 400
|
||||
assert "protocolVersion" in detail
|
||||
case _:
|
||||
pytest.fail(f"expected 400 for unsupported protocolVersion, got {result}")
|
||||
|
||||
@pytest.mark.covers("other.a2a.register.malformed_version_rejected")
|
||||
def test_malformed_protocol_version_rejected(self, client: A2AClient) -> None:
|
||||
result = _register_rejection(client, "0.3.garbage")
|
||||
match result:
|
||||
case UnknownApiError(status_code=status, body=detail):
|
||||
assert status == 400
|
||||
assert "Unsupported protocolVersion '0.3.garbage'" in detail
|
||||
case _:
|
||||
pytest.fail(f"expected 400 for malformed protocolVersion, got {result}")
|
||||
|
|
@ -26,16 +26,24 @@ from e2e_http import (
|
|||
)
|
||||
from models import LiteLLMParamsBody
|
||||
|
||||
UPLOAD_FILENAME = "batch_input.jsonl"
|
||||
|
||||
|
||||
class FileObject(BaseModel):
|
||||
id: str
|
||||
object: str | None = None
|
||||
purpose: str | None = None
|
||||
filename: str | None = None
|
||||
bytes: int | None = None
|
||||
status: str | None = None
|
||||
created_at: int | None = None
|
||||
|
||||
|
||||
class FileList(BaseModel):
|
||||
object: str | None = None
|
||||
data: list[FileObject] = []
|
||||
|
||||
|
||||
class BatchObject(BaseModel):
|
||||
id: str
|
||||
object: str | None = None
|
||||
|
|
@ -106,12 +114,30 @@ class BatchClient:
|
|||
_files_path(provider),
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
form=form,
|
||||
filename="batch_input.jsonl",
|
||||
filename=UPLOAD_FILENAME,
|
||||
content=content,
|
||||
params=ModelQuery(model=model),
|
||||
response_type=FileObject,
|
||||
)
|
||||
|
||||
def retrieve_file(
|
||||
self, file_id: str, *, key: str, provider: str | None = None
|
||||
) -> Result[FileObject]:
|
||||
return self.proxy.transport.get(
|
||||
f"{_files_path(provider)}/{file_id}",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
params=NoBody(),
|
||||
response_type=FileObject,
|
||||
)
|
||||
|
||||
def list_files(self, *, key: str, provider: str | None = None) -> Result[FileList]:
|
||||
return self.proxy.transport.get(
|
||||
_files_path(provider),
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
params=NoBody(),
|
||||
response_type=FileList,
|
||||
)
|
||||
|
||||
def create_batch(
|
||||
self, *, body: BatchCreateBody, key: str, provider: str | None = None
|
||||
) -> StreamingResponse:
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue