Merge remote-tracking branch 'origin/litellm_internal_staging' into feature/ovalix-extended-guardrail

This commit is contained in:
Shalom Jamil 2026-07-23 11:25:23 +03:00
commit 842525f2ba
508 changed files with 37534 additions and 10221 deletions

View file

@ -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:

View file

@ -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

View file

@ -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...

View file

@ -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" -->

View file

@ -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 \

View file

@ -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

View 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

View file

@ -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

View file

@ -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

View 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

View file

@ -18,6 +18,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/team/",
"/v2/team/",
"/organization/",
"/v2/organization/",
"/customer/",
"/end_user/",
"/sso/",

View file

@ -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"]

View file

@ -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")
);

View file

@ -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

View file

@ -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==",

View file

@ -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()
})
}

View file

@ -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 => {

View file

@ -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");

View file

@ -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"),

View file

@ -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!(

View file

@ -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

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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.

View file

@ -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

View file

@ -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)

View file

@ -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():

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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"

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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:

View file

@ -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

View file

@ -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
)

View file

@ -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)

View file

@ -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]

View file

@ -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

View file

@ -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"

View file

@ -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=(

View file

@ -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"))

View file

@ -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,

View file

@ -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(

View file

@ -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

View file

@ -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, ...]:

View file

@ -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)

View file

@ -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}

View file

@ -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

View file

@ -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())

View 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"]

View 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

View 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),
)

View 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)

View file

@ -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(

View file

@ -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)

View file

@ -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.

View file

@ -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

View file

@ -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(

View file

@ -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)}")

View file

@ -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()

View file

@ -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(

View file

@ -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"],

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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,
},
)

View file

@ -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,

View file

@ -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",

View file

@ -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(
{

View file

@ -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

View file

@ -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(

View file

@ -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)

View file

@ -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)

View file

@ -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:

View file

@ -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]],

View file

@ -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,

View file

@ -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]):

View file

@ -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:

View file

@ -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(

View file

@ -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]

View file

@ -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)

View file

@ -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,

View file

@ -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",
]

View file

@ -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,

View file

@ -93,7 +93,7 @@
"limit": 33
},
"DTZ005": {
"limit": 244
"limit": 241
},
"DTZ006": {
"limit": 13

View file

@ -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

View file

@ -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

View 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())

View file

@ -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.
]

View file

@ -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
View 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
View 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)

View 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}")

View file

@ -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