Merge branch 'litellm_internal_staging' into litellm_lit4658_oauth_discovery_logging

This commit is contained in:
tin-berri 2026-07-23 09:39:11 -07:00 • committed by GitHub
commit 55b0046089
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
234 changed files with 18709 additions and 3904 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

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

View file

@ -19,7 +19,7 @@ concurrency:
jobs:
ui-unit-tests:
runs-on: ubuntu-latest
runs-on: ubuntu-latest-16-cores
timeout-minutes: 20
defaults:
run:
@ -50,8 +50,8 @@ jobs:
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=4
--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=4
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

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

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

@ -9,9 +9,22 @@ duration_in_seconds is used in diff parts of the code base, example
import re
import time as time_module
from datetime import datetime, time, timedelta, timezone, tzinfo
from typing import Optional, Tuple
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
@ -124,9 +137,13 @@ def get_next_standardized_reset_time(
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

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

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

@ -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,
@ -29,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_finish_bridge_mint,
_prepare_bridge_mint,
_prepare_bridge_refresh,
_reload_active_user_by_id,
)
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
@ -39,6 +44,14 @@ from litellm.proxy._experimental.mcp_server.faults import (
dcr_fault_detail,
render_token_fault,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
aggregate_authorize,
aggregate_token,
complete_connect_flow,
is_gateway_dcr_client_id,
register_aggregate_client,
relative_request_url,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
@ -111,6 +124,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 +140,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 +164,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,14 +246,112 @@ 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 _session_cookie_user_id(request: Request) -> str | None:
"""The signed-in litellm user for a browser request, or ``None``. Thin wrapper so the
aggregate DCR flow's verbs receive the identity as a plain value instead of parsing
cookies themselves."""
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load
_user_id_from_session_cookie,
)
return _user_id_from_session_cookie(request)
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,
so a session is required; without one there is nothing to bind. After login the user re-initiates
the connection, which then finds the session cookie (the seamless return-to round-trip, which is
origin-validated against the control-plane URL, is a follow-up)."""
so a session is required; without one there is nothing to bind. A same-origin relative
``return_to`` (honored by the SSO callback) brings the browser straight back to this authorize
request after login instead of stranding it on the dashboard."""
base_url = get_request_base_url(request)
return RedirectResponse(f"{base_url}/sso/key/generate")
return RedirectResponse(f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}")
# LIT-4197: some upstream authorization servers reject an over-long ``state``
@ -623,6 +750,7 @@ 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:
@ -642,7 +770,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,
@ -686,7 +817,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)
@ -733,6 +869,7 @@ 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"):
@ -749,15 +886,24 @@ async def exchange_token_with_server(
),
)
# 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,
)
@ -1260,7 +1406,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
@ -1272,6 +1418,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,
@ -1334,30 +1616,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()
@ -1390,6 +1653,18 @@ async def authorize(
global_mcp_server_manager,
)
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
return aggregate_authorize(
request=request,
client_id=client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
response_type=response_type,
session_user_id=_session_cookie_user_id(request),
)
lookup_name: Optional[str] = mcp_server_name or client_id
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = (
@ -1453,6 +1728,25 @@ async def token_endpoint(
global_mcp_server_manager,
)
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
master_key,
user_api_key_cache,
)
return await aggregate_token(
request=request,
grant_type=grant_type,
code=code,
redirect_uri=redirect_uri,
client_id=client_id,
code_verifier=code_verifier,
refresh_token=refresh_token,
master_key=master_key,
reload_user=_reload_active_user_by_id,
cache=user_api_key_cache,
)
lookup_name = mcp_server_name or client_id
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip)
@ -1474,6 +1768,21 @@ async def token_endpoint(
)
@router.post("/authorize/complete")
async def authorize_complete(request: Request, flow: str = Form(...)):
"""Finish an aggregate connect flow: mint the gateway authorization code for the
signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly
cookie set at /authorize; an anonymous or bad-flow request just 400s."""
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load
return await complete_connect_flow(
request=request,
flow_handle=flow,
session_user_id=_session_cookie_user_id(request),
cache=user_api_key_cache,
)
# Per RFC 6749 §4.1.2.1, an IdP that rejects an OAuth authorization request
# redirects back to the configured redirect URI with ``error`` /
# ``error_description`` / ``error_uri`` query params and no ``code``. The MCP
@ -1595,11 +1904,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)
@ -2190,7 +2511,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",
@ -2199,6 +2520,13 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
}
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
# A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR
# endpoint the aggregate authorization-server metadata advertises. A single-server
# deployment registers at /{server}/register instead (its bare-origin discovery
# advertises that), so this does not affect it. A request without redirect_uris is not
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
if data.get("redirect_uris"):
return await register_aggregate_client(request=request, request_body=data)
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(

View file

@ -0,0 +1,637 @@
"""The gateway-level DCR flow for the aggregate ``/mcp`` endpoint (``mcp_gateway_dcr``).
An OAuth-only DCR client (Claude Desktop, Claude Code, MCP Inspector) pointed at the
aggregate ``/mcp`` endpoint discovers the gateway as its authorization server (PR 1 of
this track) and then walks the flow implemented here:
1. ``POST /register``: stateless dynamic client registration. The ``client_id`` IS the
registration: the client's redirect URIs are sealed into it with the repo's
authenticated symmetric helper, so nothing is persisted and a forged or tampered
client_id simply fails to open. Clients are always public (``token_endpoint_auth_method
"none"``); PKCE S256 is what protects the code.
2. ``GET /authorize``: validates the client and redirect URI, requires S256 PKCE, and
interposes LiteLLM sign-in. Without a session cookie the browser is sent through
``/sso/key/generate`` with a same-origin ``return_to`` so it lands back here after
login. With a session, the flow parameters and the SSO user are sealed into a per-flow
HttpOnly cookie (the same pattern as the upstream OAuth state relay) and the browser is
sent to the connect page, where the user authorizes individual servers (vaulting those
tokens server-side) before finishing.
3. ``POST /authorize/complete``: the deliberate finish step. A POST (not GET) bound to the
SameSite=Lax flow cookie, so a cross-site link cannot silently mint a code with the
victim's session, and the signed-in user must match the user sealed into the flow.
Mints a short-lived, single-use, gateway-sealed authorization code and redirects to the
client's registered redirect URI.
4. ``POST /token``: exchanges the code (PKCE-verified, client- and redirect-bound,
single-use) for the identity-only session tokens of
:mod:`.outbound_credentials.session_token`, re-validating that the litellm user is
still active first; the ``refresh_token`` grant rotates the pair the same way.
Nothing here stores state server-side except the single-use code guard (a TTL cache
entry). Every sealed value is authenticated encryption over the proxy salt/master key
family, opened totally (bad input maps to an OAuth error, never a raise), and every
identity is a stable reference re-validated live at mint, refresh, and (in the admission
PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they
are vaulted per user by the existing ``/v1/mcp`` authorize endpoints and resolved at
egress by user id.
"""
from __future__ import annotations
import hashlib
import hmac
import secrets
from base64 import urlsafe_b64encode
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Awaitable, Callable, Literal, TypeVar
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
is_loopback_redirect_host,
validate_redirect_uri_shape,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
SessionRefreshOpened,
open_session_refresh_bearer,
session_keys_from_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
SESSION_REFRESH_TTL_SECONDS,
MintedSessionToken,
SessionKeys,
SessionPrincipal,
mint_session_refresh_token,
mint_session_token,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
GATEWAY_DCR_CLIENT_ID_PREFIX = "llm_dcrc_"
"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token
endpoints can route an aggregate-flow request without decrypting, and existing per-server
flows (whose client_ids are upstream-issued) are never captured by the aggregate arm."""
GATEWAY_AUTH_CODE_PREFIX = "llm_gcode_"
"""Marker prefix on the gateway-sealed authorization code, distinct from the bridge
``llm_bcode_`` so neither flow can consume the other's codes."""
CONNECT_FLOW_COOKIE_PREFIX = "mcp_connect_flow_"
"""Per-flow HttpOnly cookie holding the sealed connect flow, keyed by a short random
handle carried in the connect-page URL (the same handle-plus-cookie pattern as the
``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no
server-side session store, and the sealed value never appears in a URL)."""
CONNECT_FLOW_TTL_SECONDS = 600
GATEWAY_AUTH_CODE_TTL_SECONDS = 120
_CLAIM_TTL_BUFFER_SECONDS = 60
_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:"
_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:"
_USED_REFRESH_CACHE_PREFIX = "mcp_gateway_dcr_refresh_used:"
MAX_REDIRECT_URIS = 3
MAX_REDIRECT_URI_LENGTH = 256
MAX_CLIENT_ID_LENGTH = 2048
"""Registration bounds. They exist to bound the sealed client_id, which rides inside
every session-token claim set: 3 URIs of 256 bytes seal to roughly 1.2KB, comfortably
under this cap and under the session token's own 4KB ceiling. Claude Desktop and MCP
Inspector register one or two redirect URIs."""
MAX_STATE_LENGTH = 1024
"""Bound on the client ``state`` sealed into the flow cookie and echoed on the auth-code
redirect. An unbounded ``state`` can push the sealed cookie past the browser's ~4KB cap
(silently dropped, breaking the flow); spec clients send a short opaque value."""
MIN_CODE_VERIFIER_LENGTH = 43
MAX_CODE_VERIFIER_LENGTH = 128
"""RFC 7636 section 4.1 bounds for the PKCE ``code_verifier``. Enforced so an out-of-range
verifier gets a clean ``invalid_request`` instead of an opaque PKCE-mismatch."""
_UNPREFIXED = ""
"""Prefix for a sealed value that carries no wire marker because it is never routed by
prefix (the connect flow lives only in its own per-handle cookie, opened by that one
handle). Named so the empty-string argument to ``_seal`` / ``_open_sealed`` reads as
deliberate rather than a typo."""
_CLIENT_RECORD_DEBUG_KEY = "gateway_dcr_client"
_CONNECT_FLOW_DEBUG_KEY = "gateway_connect_flow"
_AUTH_CODE_DEBUG_KEY = "gateway_authorization_code"
ReloadUserFailure = Literal["unresolvable", "unavailable", "no_active_key"]
ReloadUser = Callable[[str], Awaitable[ReloadUserFailure | None]]
"""Injected live-user revalidation (the token endpoint's mirror of admission):
``None`` means the user is active; ``unavailable`` is a retryable DB outage; anything
else fails the grant closed."""
class GatewayDcrClient(BaseModel):
"""The registration record sealed into a gateway DCR ``client_id``.
``extra="forbid"`` so a sealed value of another type (an auth code, a connect flow)
that happened to decrypt under the shared key can never validate as a client record:
cross-type confusion is rejected at the model boundary, not left to differing required
fields."""
model_config = ConfigDict(frozen=True, extra="forbid")
redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS)
iat: int
class _ConnectFlow(BaseModel):
"""One in-flight authorize: the SSO user it belongs to and the client parameters
needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti``
makes the flow single-use at complete; ``extra="forbid"`` rejects cross-type
confusion."""
model_config = ConfigDict(frozen=True, extra="forbid")
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
state: str
code_challenge: str = Field(min_length=1)
jti: str = Field(min_length=1)
exp: int
class _GatewayAuthCode(BaseModel):
"""The gateway-sealed authorization code: the user consent it represents and the
bindings the token endpoint must verify (client, redirect URI, PKCE challenge),
plus a ``jti`` for the single-use guard. ``extra="forbid"`` rejects cross-type
confusion."""
model_config = ConfigDict(frozen=True, extra="forbid")
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
code_challenge: str = Field(min_length=1)
jti: str = Field(min_length=1)
iat: int
exp: int
def is_gateway_dcr_client_id(client_id: str | None) -> bool:
"""Cheap prefix routing test so the root endpoints only enter the aggregate arm for
clients this flow registered; every other client_id keeps today's behavior."""
return client_id is not None and client_id.startswith(GATEWAY_DCR_CLIENT_ID_PREFIX)
def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse:
"""RFC 6749 section 5.2 / RFC 7591 section 3.2.2 error body. Descriptions carry no
token, code, or URL material so they are safe to relay to any client."""
return JSONResponse(
status_code=status_code,
content={"error": error, "error_description": description},
headers=TOKEN_NO_CACHE_HEADERS,
)
def _seal(prefix: str, payload: BaseModel) -> str:
return prefix + encrypt_value_helper(payload.model_dump_json())
_SealedModelT = TypeVar("_SealedModelT", bound=BaseModel)
def _open_sealed(value: str, prefix: str, model: type[_SealedModelT], debug_key: str) -> _SealedModelT | None:
"""Open a sealed value totally: anything that is not prefix-shaped, does not decrypt,
or does not validate returns ``None`` for the caller to map onto an OAuth error."""
if not value.startswith(prefix):
return None
decrypted = decrypt_value_helper(value[len(prefix) :], debug_key, return_original_value=False)
if not isinstance(decrypted, str):
return None
try:
return model.model_validate_json(decrypted)
except ValidationError:
return None
def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None:
return _open_sealed(client_id, GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient, _CLIENT_RECORD_DEBUG_KEY)
async def register_aggregate_client(request: Request, request_body: Mapping[str, object]) -> Response:
"""RFC 7591 dynamic registration against the gateway itself, statelessly.
Only ``redirect_uris`` is authoritative; every client is registered as a public
``token_endpoint_auth_method "none"`` client regardless of what it asked for (RFC
7591 lets the server override metadata), because the gateway never issues client
secrets: possession of a secret would add nothing over the mandatory S256 PKCE, and a
stateless registration has nowhere to keep one. Nothing is persisted, so open
registration cannot be used to fill storage.
Redirect-URI *hygiene* is not decided here: :func:`validate_redirect_uri_shape` is
the single owner of that rule across the MCP OAuth surface, so allowlisted native
callbacks (``cursor://``) are accepted and fragments, missing hosts, userinfo
(``https://claude.ai@attacker.example/cb``) and backslash hosts are rejected exactly
as they are on /authorize and /callback.
What this endpoint does decide is its own trust policy, which is deliberately wider
than :func:`validate_trusted_redirect_uri`'s: registration is *public*, so any https
client may register (that is what lets a hosted MCP client register at all), and the
controls are mandatory S256 PKCE plus the consent screen showing the client origin.
http is confined to loopback per RFC 8252 section 7.3.
"""
raw_uris = request_body.get("redirect_uris")
if not isinstance(raw_uris, list) or not raw_uris or len(raw_uris) > MAX_REDIRECT_URIS:
return _oauth_error(
400,
"invalid_redirect_uri",
f"redirect_uris must be a list of 1 to {MAX_REDIRECT_URIS} URIs",
)
if not all(isinstance(uri, str) and len(uri) <= MAX_REDIRECT_URI_LENGTH for uri in raw_uris):
return _oauth_error(
400,
"invalid_redirect_uri",
f"each redirect URI must be a string of at most {MAX_REDIRECT_URI_LENGTH} characters",
)
for uri in raw_uris:
parsed = urlparse(uri)
try:
if validate_redirect_uri_shape(parsed):
continue # allowlisted native callback, e.g. cursor://
except HTTPException as exc:
# The shared validator speaks HTTP; RFC 7591 registration answers with an OAuth
# error object, so translate the shape without re-deciding the rule.
return _oauth_error(400, "invalid_redirect_uri", str(exc.detail))
if parsed.scheme == "https" or (parsed.scheme == "http" and is_loopback_redirect_host(parsed)):
continue
return _oauth_error(
400,
"invalid_redirect_uri",
"each redirect URI must be https, http on a loopback host, or a registered native callback",
)
now = datetime.now(timezone.utc)
client_id = _seal(
GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient(redirect_uris=tuple(raw_uris), iat=int(now.timestamp()))
)
if len(client_id) > MAX_CLIENT_ID_LENGTH:
return _oauth_error(400, "invalid_client_metadata", "registered metadata is too large")
return JSONResponse(
status_code=201,
content={
"client_id": client_id,
"client_id_issued_at": int(now.timestamp()),
"redirect_uris": list(raw_uris),
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
},
)
def _flow_cookie_name(handle: str) -> str:
return f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}"
def _cookie_path_and_secure(request: Request) -> tuple[str, bool]:
parsed = urlparse(get_request_base_url(request))
return parsed.path or "/", parsed.scheme == "https"
def _append_query_params(url: str, params: dict[str, str]) -> str:
parsed = urlparse(url)
query = parse_qsl(parsed.query, keep_blank_values=True) + list(params.items())
return urlunparse(parsed._replace(query=urlencode(query)))
def relative_request_url(request: Request) -> str:
"""The request's own path and query as a same-origin ``return_to`` target for the
login round-trip; relative by construction, so it can never leave the gateway."""
path = request.url.path
return f"{path}?{request.url.query}" if request.url.query else path
def aggregate_authorize(
request: Request,
client_id: str,
redirect_uri: str,
state: str,
code_challenge: str | None,
code_challenge_method: str | None,
response_type: str | None,
session_user_id: str | None,
) -> Response:
"""The aggregate authorize verb: validate the client, require S256 PKCE, interpose
LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a
per-flow cookie.
Validation failures respond directly with 400 and never redirect: per RFC 6749
section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and
once the client is at fault there is no trusted place to send the browser.
"""
client = open_gateway_dcr_client(client_id)
if client is None:
return _oauth_error(400, "invalid_client", "unknown or malformed client_id")
if redirect_uri not in client.redirect_uris:
return _oauth_error(400, "invalid_request", "redirect_uri is not registered for this client")
if response_type != "code":
return _oauth_error(400, "unsupported_response_type", "response_type must be 'code'")
if not code_challenge or code_challenge_method != "S256":
return _oauth_error(
400,
"invalid_request",
"PKCE is required: send code_challenge with code_challenge_method=S256",
)
if len(state) > MAX_STATE_LENGTH:
return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters")
base_url = get_request_base_url(request)
if session_user_id is None:
login_url = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}"
return RedirectResponse(login_url, status_code=303)
now = datetime.now(timezone.utc)
handle = secrets.token_urlsafe(24)
flow = _ConnectFlow(
user_id=session_user_id,
client_id=client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
jti=secrets.token_urlsafe(24),
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
)
connect_url = _append_query_params(
f"{base_url}/ui/chat/integrations",
{"connect_flow": handle, "connect_client": _origin_only(redirect_uri)},
)
response = RedirectResponse(connect_url, status_code=303)
path, secure = _cookie_path_and_secure(request)
response.set_cookie(
key=_flow_cookie_name(handle),
value=_seal(_UNPREFIXED, flow),
max_age=CONNECT_FLOW_TTL_SECONDS,
path=path,
secure=secure,
httponly=True,
samesite="lax",
)
return response
def _origin_only(url: str) -> str:
"""Scheme+host for display on the connect page; never the full redirect URI, whose
path or query could carry values that do not belong in a page URL or logs."""
parsed = urlparse(url)
return f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else ""
async def complete_connect_flow(
request: Request,
flow_handle: str,
session_user_id: str | None,
cache: DualCache,
) -> Response:
"""The deliberate finish step of the connect flow: mint the gateway authorization
code and send the browser back to the client.
Reached by POST so a cross-site GET cannot trigger it, and bound to the HttpOnly
per-flow cookie plus an exact match between the signed-in user and the user sealed
into the flow: a link crafted by another party dies here with ``access_denied``
instead of minting a code for the victim's identity. The flow is single-use (an atomic
claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in.
"""
sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle))
if sealed_flow is None:
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
flow = _open_sealed(sealed_flow, _UNPREFIXED, _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY)
if flow is None:
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
now = datetime.now(timezone.utc)
if now.timestamp() >= flow.exp:
return _oauth_error(400, "invalid_request", "the connect flow has expired; restart the connection")
if session_user_id is None:
return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting")
if session_user_id != flow.user_id:
return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow")
if not await _SingleUseGuard(cache).claim(
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection")
code = _seal(
GATEWAY_AUTH_CODE_PREFIX,
_GatewayAuthCode(
user_id=flow.user_id,
client_id=flow.client_id,
redirect_uri=flow.redirect_uri,
code_challenge=flow.code_challenge,
jti=secrets.token_urlsafe(24),
iat=int(now.timestamp()),
exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS,
),
)
params = {"code": code, **({"state": flow.state} if flow.state else {})}
response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303)
path, secure = _cookie_path_and_secure(request)
response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax")
return response
def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool:
"""RFC 7636 S256 verification, total over hostile input. The comparison is over bytes
so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's
authorize request) simply fails to match instead of raising ``TypeError`` the way
``hmac.compare_digest`` does on two ``str`` with non-ASCII content. The verifier is
ASCII per spec; a compliant client's challenge is base64url and matches."""
digest = hashlib.sha256(code_verifier.encode("ascii", "replace")).digest()
computed = urlsafe_b64encode(digest).rstrip(b"=")
return hmac.compare_digest(computed, code_challenge.encode("utf-8"))
class _SingleUseGuard:
"""Atomic single-use claim for a one-time id (an auth-code, connect-flow ``jti``, or refresh-token
``jti``) over the injected proxy cache.
Uses an atomic increment rather than a get-then-set: two concurrent redemptions of the same id
cannot both observe "unused", because exactly one increment returns 1. The claim IS the gate, so it
fails closed. Crucially, the increment must be recorded in a backend SHARED across replicas, or the
single-use property is per-worker only (each replica's in-memory counter returns 1, so a captured
id replays through a different worker):
- When a Redis backend is configured it is the SOLE authority: the claim goes straight to Redis
(``INCR`` is atomic across replicas), and any Redis fault fails the claim CLOSED — it never falls
back to the per-worker in-memory count (``DualCache.async_increment_cache`` does fall back, which
is exactly the replay window this avoids).
- With no Redis configured (single-replica) the in-memory increment is authoritative within the one
process. A multi-worker deployment must run Redis for the guarantee to hold across workers.
The id's own TTL is the outer bound. For the auth code, PKCE binding is the primary defense against
interception; this makes the RFC 6749 4.1.2 single-use property reliable on top of it."""
def __init__(self, cache: DualCache) -> None:
self._cache = cache
async def claim(self, key: str, ttl_seconds: int) -> bool:
"""Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1); ``False``
on a replay (>1) or when the claim could not be recorded in the shared backend (fail closed)."""
from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load
# Resolve the shared authority HERE rather than trusting the injected cache: callers pass
# user_api_key_cache, which only carries a redis_cache when enable_redis_auth_cache is set
# (off by default), so a guard that read its injected cache silently degraded every claim to
# a per-worker count on a stock multi-worker deployment. redis_usage_cache is the store the
# proxy already treats as cross-worker, so no call site can wire the guarantee away.
redis_cache = redis_usage_cache or getattr(self._cache, "redis_cache", None)
if redis_cache is not None:
# Shared, atomic authority for multi-replica deployments. Claim ONLY against Redis and fail
# CLOSED on any Redis fault (async_increment re-raises) rather than fall back to the
# per-worker in-memory count, which would let each replica observe count==1 and replay the id.
try:
count = await redis_cache.async_increment(key, 1, ttl=ttl_seconds)
except Exception as e: # noqa: BLE001 # ANY Redis fault fails the single-use claim closed
verbose_logger.warning(
"mcp gateway single-use claim: shared cache backend unavailable, failing closed: %s", e
)
return False
return count == 1
# No shared backend configured (single-replica): the in-memory increment is authoritative.
count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True)
return count == 1
def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response:
access = mint_session_token(principal, keys, now)
refresh = mint_session_refresh_token(principal, keys, now)
if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken):
return _oauth_error(500, "server_error", "failed to mint the session credential")
return JSONResponse(
status_code=200,
content={
"access_token": access.token.get_secret_value(),
"token_type": "Bearer",
"expires_in": int((access.expires_at - now).total_seconds()),
"refresh_token": refresh.token.get_secret_value(),
},
headers=TOKEN_NO_CACHE_HEADERS,
)
def _reload_failure_response(failure: ReloadUserFailure) -> Response:
"""Map the live-user revalidation failure onto its OAuth error, exhaustively, so a new
``ReloadUserFailure`` member is a type error here rather than silently 400ing."""
match failure:
case "unavailable":
return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry")
case "unresolvable":
return _oauth_error(500, "server_error", "the gateway is not configured to resolve users")
case "no_active_key":
return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active")
case _:
assert_never(failure)
async def aggregate_token(
request: Request,
grant_type: str,
code: str | None,
redirect_uri: str | None,
client_id: str,
code_verifier: str | None,
refresh_token: str | None,
master_key: str | None,
reload_user: ReloadUser,
cache: DualCache,
) -> Response:
"""The aggregate token verb: authorization_code and refresh_token grants for the
identity-only session pair. Every path re-validates the litellm user live before
minting, so a deactivated user cannot obtain or renew a session."""
if master_key is None:
verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured")
return _oauth_error(500, "server_error", "the gateway has no master key configured")
keys = session_keys_from_master_key(master_key)
now = datetime.now(timezone.utc)
if grant_type == "authorization_code":
return await _authorization_code_grant(
code=code,
redirect_uri=redirect_uri,
client_id=client_id,
code_verifier=code_verifier,
keys=keys,
now=now,
reload_user=reload_user,
guard=_SingleUseGuard(cache),
)
if grant_type == "refresh_token":
return await _refresh_token_grant(
refresh_token=refresh_token,
client_id=client_id,
keys=keys,
now=now,
reload_user=reload_user,
guard=_SingleUseGuard(cache),
)
return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token")
async def _authorization_code_grant(
code: str | None,
redirect_uri: str | None,
client_id: str,
code_verifier: str | None,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
guard: _SingleUseGuard,
) -> Response:
if not code or not redirect_uri or not code_verifier:
return _oauth_error(400, "invalid_request", "code, redirect_uri, and code_verifier are required")
if not MIN_CODE_VERIFIER_LENGTH <= len(code_verifier) <= MAX_CODE_VERIFIER_LENGTH:
return _oauth_error(400, "invalid_request", "code_verifier must be 43 to 128 characters (RFC 7636)")
parsed = _open_sealed(code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY)
if parsed is None:
return _oauth_error(400, "invalid_grant", "the authorization code is invalid")
if now.timestamp() >= parsed.exp:
return _oauth_error(400, "invalid_grant", "the authorization code has expired")
if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri:
return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client")
if not _pkce_verifier_matches(code_verifier, parsed.code_challenge):
return _oauth_error(400, "invalid_grant", "PKCE verification failed")
# Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable
# 503) does not consume a still-valid code and force the client to restart sign-in.
failure = await reload_user(parsed.user_id)
if failure is not None:
return _reload_failure_response(failure)
# Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller
# wins, and a claim that cannot be recorded fails closed.
if not await guard.claim(
f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_grant", "the authorization code was already used")
return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now)
async def _refresh_token_grant(
refresh_token: str | None,
client_id: str,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
guard: _SingleUseGuard,
) -> Response:
if not refresh_token:
return _oauth_error(400, "invalid_request", "refresh_token is required")
opened = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id)
if not isinstance(opened, SessionRefreshOpened):
return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client")
failure = await reload_user(opened.principal.user_id)
if failure is not None:
return _reload_failure_response(failure)
# Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is
# single-use. Claim its jti before issuing the replacement pair, so a captured or replayed
# refresh token cannot mint a second pair after the legitimate holder rotated. Claimed AFTER
# user revalidation so a transient DB 503 does not burn a still-valid token; a claim that
# cannot be recorded fails closed, exactly like the authorization-code path.
if not await guard.claim(
f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_grant", "the refresh token was already used")
return _session_token_pair(opened.principal, keys, now)

View file

@ -50,6 +50,7 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
@ -2318,6 +2319,56 @@ class MCPServerManager:
return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None]
async def operator_open_server_ids(
self,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
allow_all_server_ids: list[str] | None = None,
submitted_server_ids: list[str] | None = None,
) -> set:
"""Servers reachable through OPEN channels rather than a grant: operator-opened
``allow_all_keys`` servers, plus the caller's own active BYOM submissions when the caller
carries no explicit ``mcp_servers`` scope.
The single owner of that question for BOTH axes. The server union in
``get_allowed_mcp_servers`` adds these ids, and the admitted subject's tool resolution asks
the same question to treat an open-channel server as default-open for tools — exactly how a
virtual key experiences it. Encoding the channel membership twice is how a server ends up
listable but uninvokable.
Empty inside a toolset scope: toolset_mcp_route / dynamic_mcp_route set
``_mcp_active_toolset_id`` before calling the handler, pinning the request to the toolset's
own servers (checking op.mcp_toolsets==[] instead would false-positive on DB-default rows
where Postgres initialises the column to ARRAY[]::TEXT[]).
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
which precomputes both for its fallback path, does not compute them twice."""
from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415
_mcp_active_toolset_id,
)
if _mcp_active_toolset_id.get() is not None:
return set()
if allow_all_server_ids is None:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
open_ids = set(allow_all_server_ids)
key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None
# "Explicitly scoped, so do not widen with BYOM" is a rule about a CREDENTIAL that carries
# its own mcp_servers list. It does not describe a keyless admitted subject: its
# object_permission is the user's own row, whose mcp_servers column is [] by DB default, so
# applying this rule would hide almost every admitted user's OWN submitted servers. Their
# submissions are theirs by authorship, and their scope comes from the per-source union.
has_explicit_object_permission = (
not _is_mcp_admitted_user_subject(user_api_key_auth)
and key_object_permission is not None
and (key_object_permission.mcp_servers is not None)
)
if not has_explicit_object_permission:
if submitted_server_ids is None:
submitted_server_ids = await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
open_ids.update(submitted_server_ids)
return open_ids
async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> list[str]:
"""
Get the allowed MCP Servers for the user.
@ -2331,11 +2382,22 @@ class MCPServerManager:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
# absolute for a scoped KEY credential are not absolute for it: its own opt-out silences its
# own source (handled per source in the resolver), never its teams' grants, and its admin
# role does not swallow the grant model — a session bearer is a third-party client
# credential, not the dashboard, so an admin signing in through the connect flow gets their
# grants like anyone else rather than handing the client the full registry ahead of every
# per-team org ceiling.
is_admitted_subject = _is_mcp_admitted_user_subject(user_api_key_auth)
# The key explicitly opted out of every MCP server. Return zero before
# layering on allow_all_keys or submitted servers so the opt-out is absolute.
key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None
if key_object_permission is not None and (
SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or [])
if (
not is_admitted_subject
and key_object_permission is not None
and (SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []))
):
return []
@ -2355,8 +2417,14 @@ class MCPServerManager:
)
try:
# If admin but NO explicit object permission, get all servers
if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission:
# If admin but NO explicit object permission, get all servers (never for an admitted
# subject — see is_admitted_subject above)
if (
user_api_key_auth
and not is_admitted_subject
and _user_has_admin_view(user_api_key_auth)
and not has_explicit_object_permission
):
verbose_logger.debug("Admin user without explicit object_permission - returning all servers")
return list(self.get_registry().keys())
@ -2364,20 +2432,14 @@ class MCPServerManager:
allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}")
combined_servers = set(allowed_mcp_servers)
# Only skip allow_all_keys servers when the request is inside a toolset
# scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id
# before calling the handler — that ContextVar is the reliable signal.
# Using op.mcp_toolsets==[] would false-positive on DB-default rows where
# Postgres initialises the column to ARRAY[]::TEXT[].
from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415
_mcp_active_toolset_id,
combined_servers.update(
await self.operator_open_server_ids(
user_api_key_auth,
allow_all_server_ids=allow_all_server_ids,
submitted_server_ids=submitted_server_ids,
)
)
in_toolset_scope = _mcp_active_toolset_id.get() is not None
if not in_toolset_scope:
combined_servers.update(allow_all_server_ids)
combined_servers.update(submitted_server_ids)
# For anonymous callers (no user_id, no role), also surface any
# servers the operator has opted into upstream-delegated auth.
# These servers handle their own auth at the upstream level, so

View file

@ -366,8 +366,36 @@ def _parse_redirect_uri_for_validation(redirect_uri: str) -> ParseResult:
)
def _validate_trusted_http_redirect_shape(parsed: ParseResult) -> bool:
"""Return True when ``parsed`` is an allowlisted native callback (caller may return)."""
def is_loopback_redirect_host(parsed: ParseResult) -> bool:
"""True when the redirect host is loopback (RFC 8252 section 7.3).
Shared by every redirect-URI policy in the MCP OAuth surface so that none of them
hand-rolls its own host list: a literal ``("localhost", "127.0.0.1", "::1")`` tuple
silently misses the rest of 127.0.0.0/8 and IPv6-mapped forms.
"""
host = (parsed.hostname or "").lower()
if host == "localhost":
return True
try:
return ip_address(host).is_loopback
except ValueError:
return False
def validate_redirect_uri_shape(parsed: ParseResult) -> bool:
"""Validate redirect-URI *hygiene* and resolve allowlisted native callbacks.
Returns True when ``parsed`` is an allowlisted native callback (the caller may accept
it outright); returns False for http/https, leaving the trust decision to the caller;
raises for a URI that no policy should ever accept (bad scheme, fragment, missing
host, userinfo, backslash in the host).
This is deliberately separate from :func:`validate_trusted_redirect_uri`, which adds
the *first-party* trust policy (same-origin, loopback, ops allowlist) appropriate to
the proxy's own OAuth endpoints. Public dynamic-client registration accepts any https
client and relies on PKCE plus the consent screen instead, so it shares this hygiene
rule but not that trust policy.
"""
if parsed.scheme not in ("http", "https"):
if _matches_trusted_native_redirect_uri(parsed):
return True
@ -419,14 +447,8 @@ def _trusted_redirect_uri_is_allowed(
):
return True
host = (parsed.hostname or "").lower()
if host == "localhost":
if is_loopback_redirect_host(parsed):
return True
try:
if ip_address(host).is_loopback:
return True
except ValueError:
pass
if parsed.scheme == "https":
for entry in _parse_trusted_redirect_origins():
@ -545,7 +567,7 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
:func:`validate_loopback_redirect_uri`.
"""
parsed = _parse_redirect_uri_for_validation(redirect_uri)
if _validate_trusted_http_redirect_shape(parsed):
if validate_redirect_uri_shape(parsed):
return
redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc)
proxy_base = _resolve_proxy_base_for_redirect(request)

View file

@ -149,6 +149,7 @@ class SessionRefreshOpened(BaseModel):
model_config = ConfigDict(frozen=True)
tag: Literal["opened"] = "opened"
principal: SessionPrincipal
jti: str
class SessionRefreshInvalid(BaseModel):
@ -187,4 +188,4 @@ def open_session_refresh_bearer(
return SessionRefreshInvalid()
if opened.principal.client_id != expected_client_id:
return SessionRefreshInvalid()
return SessionRefreshOpened(principal=opened.principal)
return SessionRefreshOpened(principal=opened.principal, jti=opened.jti)

View file

@ -113,10 +113,12 @@ class MintedSessionToken(BaseModel):
class OpenedSessionToken(BaseModel):
"""A validated session token of either kind: the principal it was minted for."""
"""A validated session token of either kind: the principal it was minted for, plus the
``jti`` so the token endpoint can enforce single-use rotation on a refresh token."""
model_config = ConfigDict(frozen=True)
principal: SessionPrincipal
jti: str
class SessionTokenTooLarge(BaseModel):
@ -320,7 +322,9 @@ def _open(
return SessionMalformed()
if now.timestamp() >= claims.exp:
return SessionExpired()
return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id))
return OpenedSessionToken(
principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti
)
def _decode_claims(

View file

@ -958,7 +958,17 @@ if MCP_AVAILABLE:
data = await add_litellm_data_to_request(
data=body_data,
request=request,
user_api_key_dict=user_api_key_auth,
# Bill a team-derived call to the team that granted it. A keyless admitted
# subject carries no team_id, so spend skipped team updates entirely and
# charged the user's PRIMARY org — the granting team's budget never
# accumulated (so it could never begin to block) and, cross-org, the wrong
# organization was charged. This is the ACCOUNTING half; the enforcement
# half (an already-over-budget team stops granting) lives in the source gate.
# Authorization is unaffected: it ran before this, and the union is resolved
# from the untouched auth object passed to call_mcp_tool below.
user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call(
user_api_key_auth, tool_name=name
),
proxy_config=proxy_config,
)
else:

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
@ -2594,6 +2605,17 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
is_session_token: bool = False
# Server-only marker set exclusively by the MCP gateway admission path
# (_reload_admitted_user) for a keyless user-subject admitted via a gateway DCR session
# bearer or bridge envelope. Not a DB column and never populated from caller-controlled key
# metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union
# or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization.
mcp_admitted_user_subject: bool = Field(default=False, exclude=True)
# team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP
# servers through several teams at once and therefore has no single team_id for the limiter to
# key off. Server-only and stripped from validated input for the same reason as the marker
# above: a forged entry would let a caller pick which team's rpm bucket it is charged against.
mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True)
budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True)
budget_throttle_pct: Optional[float] = Field(default=None, exclude=True)
user: Optional[Any] = None # Expanded user object when expand=user is used
@ -2614,6 +2636,11 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# If values is already an instance (not a dict), return it as-is
if not isinstance(values, dict):
return values
# mcp_admitted_user_subject is a server-only marker, set ONLY by the MCP gateway admission
# path via post-construction assignment. Strip it from any validated input (constructor
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
values.pop("mcp_source_team_rpm_limits", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):
@ -2767,6 +2794,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,
)

View file

@ -2661,6 +2661,15 @@ async def get_managed_vector_store_rows_by_uuids(
return result
class OrganizationNotFoundError(Exception):
"""The organization row is CONFIRMED absent, as opposed to a lookup that failed.
Subclasses Exception so every existing except Exception caller keeps its current
behavior; it exists so a caller that wants to treat "no such org" as "no restriction" can do
that WITHOUT also swallowing an outage and silently dropping a real org ceiling.
"""
@log_db_metrics
async def get_org_object(
org_id: str,
@ -2707,25 +2716,30 @@ async def get_org_object(
query_kwargs["include"] = {"litellm_budget_table": True}
response = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs)
if response is None:
raise Exception
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_org_obj,
model_type=LiteLLM_OrganizationTable,
ttl=DEFAULT_IN_MEMORY_TTL,
)
return _org_obj
except Exception:
raise Exception(
# An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed
# missing row, and relabelling it as "doesn't exist" made every caller unable to tell them
# apart — a caller that treats absence as "this org places no restriction" then drops a real
# org ceiling during an outage. Propagate the real error; callers that already catch
# Exception are unaffected.
raise
if response is None:
raise OrganizationNotFoundError(
f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call."
)
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_org_obj,
model_type=LiteLLM_OrganizationTable,
ttl=DEFAULT_IN_MEMORY_TTL,
)
return _org_obj
async def _get_resources_from_access_groups(
access_group_ids: List[str],

View file

@ -7,12 +7,15 @@ login endpoints (e.g., /login and /v2/login).
import os
import secrets
from datetime import datetime, timedelta, timezone
from typing import Literal, Optional, cast
import jwt
from fastapi import HTTPException
import litellm
from litellm.constants import LITELLM_PROXY_ADMIN_NAME, LITELLM_UI_SESSION_DURATION
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import (
LiteLLM_UserTable,
LitellmUserRoles,
@ -313,6 +316,29 @@ async def authenticate_user(
)
def _ui_session_exp_timestamp() -> int:
"""The ``exp`` claim (unix seconds) for a UI session cookie, ``LITELLM_UI_SESSION_DURATION``
from now. The virtual key sealed inside the cookie already expires after this same
duration; stamping the JWT itself gives the cookie the bounded lifetime the dashboard's
client-side expiry check and the server-side session-cookie readers both assume, instead
of a token that stays signature-valid until the master key rotates."""
ttl_seconds = duration_in_seconds(LITELLM_UI_SESSION_DURATION)
return int((datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)).timestamp())
def encode_ui_session_jwt(returned_ui_token_object: ReturnedUITokenObject, master_key: str) -> str:
"""Encode a UI session cookie JWT with a bounded ``exp``.
The single choke point every UI login path (SSO and username/password /login, /v2,
/v3) uses to mint the ``token`` cookie, so the cookie's lifetime is set in exactly one
place and cannot drift between paths. Without the ``exp`` the cookie is valid until the
master key rotates, and the session-cookie readers that require a bounded lifetime
(the MCP interactive sign-in) reject it.
"""
claims = {**cast(dict, returned_ui_token_object), "exp": _ui_session_exp_timestamp()}
return jwt.encode(claims, master_key, algorithm="HS256")
def create_ui_token_object(
login_result: LoginResult,
general_settings: dict,

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

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

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

@ -1781,28 +1781,38 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.team_id:
if not mcp_server_name:
return
mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
# Which teams' buckets does this call charge? A key is pinned to exactly one team. A keyless
# MCP-admitted subject reaches servers through SEVERAL teams at once and has no team_id, so
# without the second source below its calls charged no team bucket at all and it outran every
# team's mcp_rpm_limit. Every applicable team is charged rather than one being picked: the
# limiter enforces all descriptors, so each team's own ceiling binds on a call made through
# its grant, and there is no arbitrary attribution when several teams grant the same server.
team_limits: list[tuple[str | None, dict[str, int] | None]] = []
if user_api_key_dict.team_id:
team_limits.append((user_api_key_dict.team_id, get_team_mcp_rpm_limit(user_api_key_dict)))
for source_team_id, source_limit in (user_api_key_dict.mcp_source_team_rpm_limits or {}).items():
team_limits.append((source_team_id, source_limit))
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{user_api_key_dict.team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
for team_id, mcp_rpm_limit in team_limits:
if not team_id or not mcp_rpm_limit:
continue
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
continue
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
)
def _should_enforce_rate_limit(
self,

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
@ -35,7 +36,7 @@ if TYPE_CHECKING:
import httpx
import jwt
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from fastapi.responses import RedirectResponse
import litellm
@ -258,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(
@ -270,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."
),
)
@ -278,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:
@ -593,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)),
)
@ -612,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 = (
(
@ -644,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")
@ -670,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)
@ -861,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,
)
@ -912,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(
@ -954,15 +965,8 @@ async def google_login(
state=cli_state,
request=request,
)
if return_to is not None and sso_redirect is not None:
if SSOAuthenticationHandler._validate_return_to(return_to):
sso_redirect.set_cookie(
key="litellm_cp_return_to",
value=return_to,
max_age=600,
httponly=True,
samesite="lax",
)
if sso_redirect is not None:
_persist_return_to_cookie(sso_redirect, return_to)
return sso_redirect
from fastapi.responses import HTMLResponse
@ -971,13 +975,19 @@ async def google_login(
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
or general_settings.get("hide_default_credentials_hint", False) is True
)
return HTMLResponse(
form_response = HTMLResponse(
content=build_ui_login_form(
show_deprecation_banner=True,
hide_default_credentials_hint=hide_default_credentials_hint,
),
status_code=200,
)
# Preserve return_to across the username/password sign-in too, via the SAME shared, never-raising
# helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the
# dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always
# renders, since the helper never raises on a bad return_to).
_persist_return_to_cookie(form_response, return_to)
return form_response
def generic_response_convertor(
@ -1957,6 +1967,7 @@ 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,
@ -2006,7 +2017,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)}"
@ -2037,13 +2048,14 @@ async def cli_sso_callback(
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)
@ -2083,6 +2095,7 @@ 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,
@ -2114,10 +2127,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")
@ -2192,7 +2205,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 = {
@ -2404,6 +2417,92 @@ async def sso_readiness():
)
def _is_same_origin_return_path(return_to: str) -> bool:
"""True for a strictly relative return path that stays on the gateway's own origin by
construction, and is therefore safe to honor without a configured ``control_plane_url``.
Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands
back on the authorize request.
Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers
fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or
whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out
of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather
than relying on downstream header encoding to neutralize it."""
if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to:
return False
return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to)
async def _sso_return_to_redirect(
return_to: str | None,
jwt_token: str,
redis_usage_cache,
user_api_key_cache,
) -> RedirectResponse | None:
"""Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard.
Two arms, both clearing the one-shot ``litellm_cp_return_to`` cookie:
- **Same-origin relative path** (the MCP gateway DCR authorize round-trip): set the session cookie
exactly like the dashboard path, then send the browser back where it came from.
- **Control-plane cross-origin** (``control_plane_url``): stash the JWT behind a single-use opaque
code (60s TTL) so the token never lands in browser history/logs; the control plane redeems it via
``POST /v3/login/exchange``.
Extracted from ``get_redirect_response_from_openid`` to keep that method inside the complexity
budget; behavior is identical to the inline arms it replaces (including letting
``_validate_return_to`` raise for a mismatched absolute return_to, as before)."""
if return_to is None:
return None
if _is_same_origin_return_path(return_to):
redirect_response = RedirectResponse(url=return_to, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
if SSOAuthenticationHandler._validate_return_to(return_to):
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": return_to}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
else:
await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
separator = "&" if "?" in return_to else "?"
redirect_url = return_to + separator + urlencode({"login": "success", "code": code})
verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code")
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
return None
def _persist_return_to_cookie(response: Response, return_to: str | None) -> None:
"""Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to``
cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there
afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot
diverge (a per-branch reimplementation is exactly how the two drifted before). Honors a strictly
relative same-origin path, and (when ``control_plane_url`` is configured) a return_to matching that
origin. It NEVER raises: a mismatched or invalid ``return_to`` is simply not stored, so it can never
block sign-in — the login entrypoint must always render."""
if return_to is None:
return
try:
safe = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to)
except HTTPException:
return # a non-matching absolute return_to is ignored, never blocks sign-in
if safe:
response.set_cookie(
key="litellm_cp_return_to",
value=return_to,
max_age=600,
httponly=True,
samesite="lax",
)
class SSOAuthenticationHandler:
"""
Handler for SSO Authentication across all SSO providers
@ -3041,7 +3140,6 @@ class SSOAuthenticationHandler:
return_to: Optional[str] = None,
sso_assertion: SSOIdentityAssertion | None = None,
) -> RedirectResponse:
import jwt
from litellm.proxy.proxy_server import (
general_settings,
@ -3205,30 +3303,21 @@ class SSOAuthenticationHandler:
server_root_path=get_server_root_path(),
)
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
master_key or "",
algorithm="HS256",
from litellm.proxy.auth.login_utils import encode_ui_session_jwt
jwt_token = encode_ui_session_jwt(returned_ui_token_object, master_key or "")
# Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane
# cross-origin code exchange) lives in one shared helper so this method stays inside the
# complexity budget. None falls through to the dashboard redirect below.
return_to_redirect = await _sso_return_to_redirect(
return_to=return_to,
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
)
# Control-plane cross-origin: store JWT behind a single-use opaque
# code (60s TTL) so the token never appears in browser history / logs.
# The control plane redeems it via POST /v3/login/exchange.
if return_to is not None and SSOAuthenticationHandler._validate_return_to(return_to):
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": return_to}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
else:
await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
separator = "&" if "?" in return_to else "?"
redirect_url = return_to + separator + urlencode({"login": "success", "code": code})
verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code")
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
if return_to_redirect is not None:
return return_to_redirect
if user_id is not None and isinstance(user_id, str):
litellm_dashboard_ui += "?login=success"

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,
@ -303,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,
@ -1158,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"}
@ -1970,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
@ -3696,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,
@ -4618,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:
@ -13451,7 +13472,7 @@ async def fallback_login(request: Request):
@router.post("/login", include_in_schema=False) # hidden since this is a helper for UI sso login
async def login(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
form = await request.form()
@ -13474,13 +13495,7 @@ async def login(request: Request):
)
# Generate JWT token
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
# Build redirect URL
litellm_dashboard_ui = get_custom_url(str(request.base_url))
@ -13490,16 +13505,51 @@ async def login(request: Request):
litellm_dashboard_ui += "/ui/"
litellm_dashboard_ui += "?login=success"
# Honor a same-origin return_to preserved by the sign-in page (e.g. the aggregate DCR connect flow's
# authorize round-trip), mirroring the SSO callback; otherwise land on the dashboard. Gated by
# _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the
# one-shot cookie is cleared after use.
from litellm.proxy.management_endpoints.ui_sso import _sso_return_to_redirect
# Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm.
# _persist_return_to_cookie stores both shapes it accepts (a relative same-origin path AND a
# control_plane_url-matching absolute URL); honoring only the relative one here silently dropped
# the control-plane case, landing the user on the dashboard. One function decides how a stored
# return_to is honored for EVERY sign-in branch, so the write and read sets cannot diverge: it
# sets the token cookie on the same-origin arm and hands off via a one-time login code on the
# cross-origin arm, and clears the one-shot cookie in both.
cp_return_to = request.cookies.get("litellm_cp_return_to")
if cp_return_to:
try:
resumed = await _sso_return_to_redirect(
return_to=cp_return_to,
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
)
except Exception: # noqa: BLE001 # resuming must NEVER block a completed sign-in
# The symmetric half of _persist_return_to_cookie's "never raises" contract. The resumer
# rejects a return_to that no longer matches control_plane_url (a config change between
# the cookie's write and this read), and the user has ALREADY authenticated here —
# failing their login over a stale one-shot cookie is the worst possible outcome. Land
# on the dashboard instead; the cookie is cleared below either way.
verbose_proxy_logger.info("Ignoring stale litellm_cp_return_to cookie; landing on dashboard")
resumed = None
if resumed is not None:
return resumed
# Create redirect response with cookie
redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
if cp_return_to:
redirect_response.delete_cookie(key="litellm_cp_return_to")
return redirect_response
@router.post("/v2/login", include_in_schema=False) # hidden helper for UI logins via API
async def login_v2(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
try:
@ -13520,13 +13570,7 @@ async def login_v2(request: Request):
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
@ -13570,7 +13614,7 @@ async def login_v2(request: Request):
) # control-plane login — always returns token in body for cross-origin use
async def login_v3(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
try:
@ -13599,13 +13643,7 @@ async def login_v3(request: Request):
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
@ -14914,10 +14952,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] = {
@ -14943,21 +14982,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):
@ -14981,6 +15031,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)
@ -15003,7 +15060,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)
@ -15178,7 +15235,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
@ -15456,14 +15513,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()
@ -15479,19 +15532,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

@ -15,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,
@ -27,16 +32,6 @@ 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
@ -109,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):
@ -757,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(
@ -765,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
@ -826,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": {},
@ -881,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
@ -924,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:
@ -975,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
}

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

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

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

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

@ -22,6 +22,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `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

View file

@ -5,9 +5,9 @@ answers or when credentials/env are missing; they never skip. Pure unit coverage
of the harness itself carries no `e2e` marker and runs regardless of whether a
proxy is up.
Lifecycle: the `resources` fixture maps the init -> run -> teardown contract
(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and
teardown deletes every resource the test created on the long-lived proxy.
Lifecycle: the `resources` fixture hands each test a lifecycle.ResourceManager -
the test registers a cleanup for every resource it creates, and the fixture's
teardown deletes them all on the long-lived proxy, even when the test fails.
Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these
shared fixtures build on it.

View file

@ -32,7 +32,7 @@
- {id: llm.files.hosted_vllm.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible file upload"}
- {id: llm.rerank.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_rerank_e2e.py:29", rationale: "Cohere rerank, top_n + relevance_score"}
- {id: llm.files.openai.content.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "GET /v1/files/{id}/content returns uploaded batch JSONL bytes"}
- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_nova_sonic_realtime_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"}
- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_realtime_bedrock_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"}
- {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"}
- {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"}
- {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"}

View file

@ -72,6 +72,8 @@ POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))
POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5"))
REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60"))
EXPECT_RUST = os.environ.get("E2E_EXPECT_RUST", "").strip().lower() in ("1", "true", "yes")
LOAD_USERS = int(os.environ.get("E2E_LOAD_USERS", "750"))
LOAD_SPAWN_RATE = float(os.environ.get("E2E_LOAD_SPAWN_RATE", "50"))
LOAD_DURATION_SECONDS = float(os.environ.get("E2E_LOAD_DURATION_SECONDS", "60"))

View file

@ -1,13 +1,11 @@
"""Lifecycle contract and resource cleanup for stateful e2e tests.
"""Resource cleanup for stateful e2e tests.
Shared by every e2e suite under tests/e2e/. The proxy under test is
long-lived and never reset between tests, so anything a test creates (keys,
customers, teams, orgs, users, guardrails, budgets, ...) persists unless
explicitly deleted. Every check follows an init -> run -> teardown lifecycle;
teardown releases each resource init() created, even when run() raises.
In pytest terms (see conftest.py): the `resources` fixture's setup is init(),
the test body is run(), and the fixture's teardown is teardown().
explicitly deleted. The `resources` fixture (see conftest.py) hands each test a
ResourceManager; the test registers a cleanup for every resource it creates, and
the fixture's teardown releases them all even when the test body raises.
"""
from dataclasses import dataclass, field
@ -17,38 +15,6 @@ from proxy_client import ProxyClient
from models import KeyGenerateBody
@runtime_checkable
class E2ECase(Protocol):
"""A stateful e2e check run against a long-lived proxy.
init() acquires resources, run() exercises behaviour and asserts, teardown()
releases everything init() created. teardown() must run even if init() fails
partway or run() raises.
"""
def init(self) -> None: ...
def run(self) -> None: ...
def teardown(self) -> None: ...
def run_case(case: E2ECase) -> None:
"""Drive a case through its lifecycle: init -> run -> teardown.
teardown always runs - even when init() fails partway or run() raises (or
skips) - so resources the case already registered on the long-lived proxy are
released. init() is inside the try because cases register cleanups
progressively (e.g. create team, then user, then key), and a failure after
the first creation must still release what came before.
"""
try:
case.init()
case.run()
finally:
case.teardown()
@runtime_checkable
class ResourceClient(Protocol):
"""Proxy operations the convenience creators use. Resource types without a

View file

@ -12,7 +12,7 @@ from __future__ import annotations
import pytest
from e2e_config import unique_marker
from e2e_config import EXPECT_RUST, unique_marker
from e2e_http import StreamingResponse, require_successful_call, unwrap
from endpoints_client import EndpointsClient
from lifecycle import ResourceManager
@ -50,6 +50,13 @@ def _assert_streamed_ok(result: StreamingResponse) -> None:
assert any("message_stop" in event for event in result.stream_events), (
"stream never reached message_stop"
)
if EXPECT_RUST:
assert result.headers.get("x-litellm-rust") == "true", (
"E2E_EXPECT_RUST is set, so this gateway must serve /v1/messages through the "
"Rust path, but the response carried no x-litellm-rust marker. The request "
"still succeeded, which is exactly the failure mode: a gateway whose native "
f"extension is unavailable falls back to Python silently. headers={result.headers}"
)
class TestAzureFoundryMessages:

View file

@ -62,7 +62,7 @@ class TestSummarizePlannedTurns:
class TestRetried:
def test_transient_failures_then_success_returns_the_success(self) -> None:
outcome = Success(data=SessionMessagesResponse())
outcome = Success[SessionMessagesResponse](status_code=200, data=SessionMessagesResponse())
calls = iter(
(NetworkError(message="overloaded"), NetworkError(message="overloaded"), outcome)
)
@ -88,7 +88,7 @@ class TestRetried:
raise AssertionError("slept after a successful attempt")
result = retried(
lambda: Success(data=SessionMessagesResponse()),
lambda: Success[SessionMessagesResponse](status_code=200, data=SessionMessagesResponse()),
attempts=3,
sleep=sleep_means_retry,
)

View file

@ -1,28 +1,36 @@
"""Live e2e: a tiny max_budget on an entity actually blocks requests.
Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
teardown() deletes everything init() created (always runs, even on failure/skip).
Covers the entities with no prior live coverage - internal user, end-user,
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
One test per budget level (key, team, internal user, end-user, organization,
team member): put the tiny cap on that level, drive spend until a
`budget_exceeded` block, and where a cap could be confused with a neighbor,
prove isolation with an uncapped control key that must keep serving. The
capped-key sweep proves the key's own max_budget blocks across mint shapes
(personal, team, team-member) with roomy surroundings, so the key-level cap is
provably the blocker no matter who the key was minted to.
A non-budget error fails hard (never a skip); if calls never get blocked, budget
enforcement is broken -> fail.
"""
import time
from dataclasses import dataclass, field
from typing import Callable, List, Type
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call
from lifecycle import run_case
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
TINY_CAP = 3e-6
ROOMY_CAP = 100.0
def _chat(client: BudgetClient, key: str, *, user: str | None = None) -> StreamingResponse:
return client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16, user=user)
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
"""Send paid calls until the entity's budget blocks one; return the blocked
response so callers can assert on its shape. Key/user/org/member block within
@ -30,13 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
enforces off table spend that lands on the batch write, so it takes a few
more. A non-budget error fails hard (never a skip)."""
for _ in range(40):
result = client.chat(
key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
user=user or None,
)
result = _chat(client, key, user=user or None)
if is_budget_block(result):
return result
require_successful_call(result)
@ -44,225 +46,154 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
pytest.fail("budget never enforced within the call budget")
@dataclass
class _BudgetCase:
"""Base E2ECase: a key under some budgeted entity must get blocked.
Subclasses set up the budgeted entity in init() and register every created id
in `_undo` (run LIFO in teardown so a key is deleted before its team/org).
"""
client: BudgetClient
key: str = ""
_undo: List[Callable[[], None]] = field(
default_factory=list
) # mutable-ok: per-case teardown registry
def init(self) -> None:
raise NotImplementedError
def run(self) -> None:
_assert_budget_blocks(self.client, self.key)
def teardown(self) -> None:
for undo in reversed(self._undo):
undo()
def _assert_blocked_429(client: BudgetClient, key: str) -> StreamingResponse:
blocked = _assert_budget_blocks(client, key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
return blocked
class KeyBudgetCase(_BudgetCase):
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
key-level budget can be the thing that blocks. The refusal must be a 429
budget_exceeded; any other error already fails via _assert_budget_blocks."""
class TestBudgetBlocksPerLevel:
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_bare_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(key))
def init(self) -> None:
self.key = self.client.generate_key(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_key(self.key))
_assert_blocked_429(client, key)
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.team.blocks_over_limit")
def test_team_budget_blocks_every_team_key(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=TINY_CAP)
resources.defer(lambda: client.delete_team(team_id))
spender_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(spender_key))
sibling_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(sibling_key))
class TeamBudgetCase(_BudgetCase):
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
a key-level budget. Key A is driven until the team cap blocks it; key B's very
first call must then be refused too, proving the cap sits on the team, not the
key that spent. Both refusals must be 429 budget_exceeded."""
def init(self) -> None:
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._sibling_key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
sibling = self.client.chat(
self._sibling_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
_assert_blocked_429(client, spender_key)
sibling = _chat(client, sibling_key)
assert is_budget_block(sibling) and sibling.status_code == 429, (
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
f"got {sibling.status_code}: {sibling.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit")
def test_user_budget_enforced_across_all_their_keys(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=TINY_CAP)
resources.defer(lambda: client.delete_user(user_id))
first_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(first_key))
second_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(second_key))
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
resources.defer(lambda: client.delete_team(team_id))
client.add_team_member(team_id, user_id)
team_key = client.generate_key(team_id=team_id, user_id=user_id)
resources.defer(lambda: client.delete_key(team_key))
class InternalUserBudgetCase(_BudgetCase):
"""A user's max_budget follows the person, not the key. The capped user holds
two personal keys (no team, no key budgets) plus a team-member key on an
uncapped team; once the first personal key is refused, the other two must be
refused as well - a second key is not a fresh allowance, and since #32005 the
user budget draws down team keys too. All refusals must be 429 budget_exceeded."""
def init(self) -> None:
user_id = self.client.create_user(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_user(user_id))
self.key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._second_key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._second_key))
team_id = self.client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
self._undo.append(lambda: self.client.delete_team(team_id))
self.client.add_team_member(team_id, user_id)
self._team_key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._team_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
for label, key in (("second personal key", self._second_key), ("team-member key", self._team_key)):
result = self.client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16)
_assert_blocked_429(client, first_key)
for label, key in (("second personal key", second_key), ("team-member key", team_key)):
result = _chat(client, key)
assert is_budget_block(result) and result.status_code == 429, (
f"the {label} of a user over budget must get the same 429 budget_exceeded, "
f"got {result.status_code}: {result.body[:200]}"
)
class EndUserBudgetCase(_BudgetCase):
def init(self) -> None:
@pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit")
def test_end_user_budget_blocks_attributed_calls(
self, client: BudgetClient, resources: ResourceManager
) -> None:
customer = f"e2e-budget-cust-{unique_marker()}"
self.client.create_customer(customer, max_budget=3e-6)
self._undo.append(lambda: self.client.delete_customers([customer]))
self.key = self.client.generate_key(models=["claude-haiku-4-5"])
self._undo.append(lambda: self.client.delete_key(self.key))
self._customer = customer
client.create_customer(customer, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_customers([customer]))
key = client.generate_key(models=["claude-haiku-4-5"])
resources.defer(lambda: client.delete_key(key))
def run(self) -> None:
_assert_budget_blocks(self.client, self.key, user=self._customer)
_assert_budget_blocks(client, key, user=customer)
@pytest.mark.covers("quota_management.budget.organization.blocks_over_limit")
def test_org_budget_blocks_keys_under_it(self, client: BudgetClient, resources: ResourceManager) -> None:
org_id = client.create_org(max_budget=TINY_CAP, alias=f"e2e-budget-org-{unique_marker()}")
resources.defer(lambda: client.delete_org(org_id))
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(key))
class OrganizationBudgetCase(_BudgetCase):
"""Org carries the tiny budget; the team under it and the key carry none, so
the org is the only entity that can block (the historically weak link). The
refusal must be a 429 budget_exceeded that names the org as the blocker."""
def init(self) -> None:
self._org_id = self.client.create_org(
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
)
self._undo.append(lambda: self.client.delete_org(self._org_id))
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", organization_id=self._org_id
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
assert f"Organization={self._org_id}" in blocked.body, (
blocked = _assert_blocked_429(client, key)
assert f"Organization={org_id}" in blocked.body, (
f"refusal must name the org as the blocker, got: {blocked.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit")
def test_member_budget_blocks_without_touching_teammates(
self, client: BudgetClient, resources: ResourceManager
) -> None:
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
member_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(member_id))
client.add_team_member(team_id, member_id, max_budget_in_team=TINY_CAP)
member_key = client.generate_key(team_id=team_id, user_id=member_id)
resources.defer(lambda: client.delete_key(member_key))
teammate_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(teammate_id))
client.add_team_member(team_id, teammate_id)
teammate_key = client.generate_key(team_id=team_id, user_id=teammate_id)
resources.defer(lambda: client.delete_key(teammate_key))
class TeamMemberBudgetCase(_BudgetCase):
"""Member A's per-team budget is tiny while the team and both members' user
budgets are roomy (100.0), so the only cap that can trip is A's: a block
proves member-level enforcement and must be a 429 budget_exceeded. Teammate
B, uncapped on the same team, must keep serving after A is cut off, proving
the member cap does not leak onto the team or its members."""
def init(self) -> None:
self._team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
)
self._undo.append(lambda: self.client.delete_team(self._team_id))
self._member_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(self._member_id))
self.client.add_team_member(self._team_id, self._member_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(team_id=self._team_id, user_id=self._member_id)
self._undo.append(lambda: self.client.delete_key(self.key))
teammate_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(teammate_id))
self.client.add_team_member(self._team_id, teammate_id)
self._teammate_key = self.client.generate_key(team_id=self._team_id, user_id=teammate_id)
self._undo.append(lambda: self.client.delete_key(self._teammate_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
teammate = self.client.chat(
self._teammate_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
require_successful_call(teammate)
_assert_blocked_429(client, member_key)
require_successful_call(_chat(client, teammate_key))
def _case_id(case_cls: Type[_BudgetCase]) -> str:
return case_cls.__name__
class TestKeyBudgetBlocksAcrossKeyKinds:
"""The tiny max_budget sits on the key itself while every budget around it
(user / team / membership) is roomy, so only the key-level cap can block; the
uncapped control key minted to the same surroundings must keep serving after
the capped key is refused, proving nothing around the key was the blocker."""
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_personal_key_blocks_over_its_own_budget(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(user_id))
capped_key = client.generate_key(user_id=user_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(control_key))
@pytest.mark.parametrize(
"case_cls",
[
pytest.param(
KeyBudgetCase,
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
),
pytest.param(
TeamBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
),
pytest.param(
InternalUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),
),
pytest.param(
EndUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit"),
),
pytest.param(
OrganizationBudgetCase,
marks=pytest.mark.covers("quota_management.budget.organization.blocks_over_limit"),
),
pytest.param(
TeamMemberBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit"),
),
],
ids=_case_id,
)
def test_budget_enforcement(
client: BudgetClient, case_cls: Type[_BudgetCase]
) -> None:
run_case(case_cls(client))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_team_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
capped_key = client.generate_key(team_id=team_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(control_key))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_team_member_key_blocks_over_its_own_budget(
self, client: BudgetClient, resources: ResourceManager
) -> None:
team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
member_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(member_id))
client.add_team_member(team_id, member_id, max_budget_in_team=ROOMY_CAP)
capped_key = client.generate_key(team_id=team_id, user_id=member_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(team_id=team_id, user_id=member_id)
resources.defer(lambda: client.delete_key(control_key))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))

View file

@ -12,6 +12,7 @@ from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
TINY_CAP = 3e-6
ROOMY_CAP = 100.0
WINDOW = "30s"
RESET_DEADLINE_SECONDS = 150
@ -46,7 +47,7 @@ def _poll_until_serves_again(client: BudgetClient, key: str) -> None:
pytest.fail(f"budget never reset within {RESET_DEADLINE_SECONDS}s")
class TestBudgetResetDiagonal:
class TestBudgetResetPerLevel:
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_bare_key_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=TINY_CAP, budget_duration=WINDOW)
@ -115,3 +116,42 @@ class TestBudgetResetDiagonal:
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
class TestKeyBudgetResetAcrossKeyKinds:
"""The tiny max_budget and its 30s window sit on the key itself while the user,
team, and membership around it are roomy (100.0), so the key's own budget is
the only thing that can block and the only thing that has to reset."""
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_personal_key_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
user_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(user_id))
key = client.generate_key(user_id=user_id, max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_team_key_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-key-reset-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id, max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_team_member_key_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-key-reset-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
member_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(member_id))
client.add_team_member(team_id, member_id, max_budget_in_team=ROOMY_CAP)
key = client.generate_key(team_id=team_id, user_id=member_id, max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)

View file

@ -23,7 +23,7 @@ import pytest
from e2e_http import Result, Success
from lifecycle import ResourceManager
from models import ChatResponse, SpendLogs, SpendLogsParams
from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
pytestmark = pytest.mark.e2e
@ -232,14 +232,12 @@ def test_cache_hit_is_zero_cost_and_suffixed(
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs)
)
cache_rows = [r for r in rows if r.cache_hit == "True"]
if not cache_rows:
pytest.skip(
"no cache-hit row observed; caching may be disabled on this proxy. "
f"rows seen: {_summarize(rows)}"
)
cache_row = cache_rows[0]
cache_row = _require_row(
rows,
lambda r: r.cache_hit == "True",
"with cache_hit=True (caching is enabled on the e2e proxy, so an identical "
"repeat call must hit the cache)",
)
assert (
cache_row.spend or 0
) == 0.0, f"cache hit was charged (double-charge regression): {_summarize(rows)}"
@ -504,22 +502,27 @@ def test_each_model_on_a_shared_key_gets_its_own_row(
@pytest.mark.covers("quota_management.spend_tracking.failure.writes_failure_row")
def test_failure_call_writes_failure_status_row(
client: SpendClient, scoped_key: str
client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
result = client.chat(scoped_key, "gemini-2.5-flash", "", max_tokens=1)
if is_ok(result):
pytest.skip("call unexpectedly succeeded; could not induce a failure row")
model = f"e2e-spend-failure-{unique_marker()}"
model_id = client.proxy.create_model(
model,
LiteLLMParamsBody(model="openai/gpt-5.5", api_key="sk-invalid-e2e-failure-row"),
)
resources.defer(lambda: client.proxy.delete_model(model_id))
result = client.chat(scoped_key, model, f"trigger failure {unique_marker()}", max_tokens=1)
assert not is_ok(result), (
f"a call to a deployment with an invalid upstream key must fail, not succeed: {result}"
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.status == "failure" for r in rs)
)
failure_rows = [r for r in rows if r.status == "failure"]
if not failure_rows:
pytest.skip(
"no failure-status row was logged for the rejected call; "
"failure logging is environment-specific"
)
assert (failure_rows[0].spend or 0) == 0.0, "failed call must not be charged"
failure_row = _require_row(
rows, lambda r: r.status == "failure", "with status=failure for the rejected call"
)
assert (failure_row.spend or 0) == 0.0, "failed call must not be charged"
@pytest.mark.covers("quota_management.spend_tracking.spend_calculate.returns_cost")

111
tests/e2e/ui/package-lock.json generated Normal file
View file

@ -0,0 +1,111 @@
{
"name": "litellm-ui-e2e",
"version": "0.0.0",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "litellm-ui-e2e",
"version": "0.0.0",
"devDependencies": {
"@playwright/test": "1.58.1",
"@types/node": "20.19.37",
"typescript": "5.9.3"
}
},
"node_modules/@playwright/test": {
"version": "1.58.1",
"resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.1.tgz",
"integrity": "sha512-6LdVIUERWxQMmUSSQi0I53GgCBYgM2RpGngCPY7hSeju+VrKjq3lvs7HpJoPbDiY5QM5EYRtRX5fvrinnMAz3w==",
"dev": true,
"license": "Apache-2.0",
"dependencies": {
"playwright": "1.58.1"
},
"bin": {
"playwright": "cli.js"
},
"engines": {
"node": ">=18"
}
},
"node_modules/@types/node": {
"version": "20.19.37",
"resolved": "https://registry.npmjs.org/@types/node/-/node-20.19.37.tgz",
"integrity": "sha512-8kzdPJ3FsNsVIurqBs7oodNnCEVbni9yUEkaHbgptDACOPW04jimGagZ51E6+lXUwJjgnBw+hyko/lkFWCldqw==",
"dev": true,
"license": "MIT",
"dependencies": {
"undici-types": "~6.21.0"
}
},
"node_modules/fsevents": {
"version": "2.3.2",
"resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz",
"integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==",
"dev": true,
"hasInstallScript": true,
"license": "MIT",
"optional": true,
"os": [
"darwin"
],
"engines": {
"node": "^8.16.0 || ^10.6.0 || >=11.0.0"
}
},
"node_modules/playwright": {
"version": "1.58.1",
"resolved": "https://registry.npmjs.org/playwright/-/playwright-1.58.1.tgz",
"integrity": "sha512-+2uTZHxSCcxjvGc5C891LrS1/NlxglGxzrC4seZiVjcYVQfUa87wBL6rTDqzGjuoWNjnBzRqKmF6zRYGMvQUaQ==",
"dev": true,
"license": "Apache-2.0",
"dependencies": {
"playwright-core": "1.58.1"
},
"bin": {
"playwright": "cli.js"
},
"engines": {
"node": ">=18"
},
"optionalDependencies": {
"fsevents": "2.3.2"
}
},
"node_modules/playwright-core": {
"version": "1.58.1",
"resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.58.1.tgz",
"integrity": "sha512-bcWzOaTxcW+VOOGBCQgnaKToLJ65d6AqfLVKEWvexyS3AS6rbXl+xdpYRMGSRBClPvyj44njOWoxjNdL/H9UNg==",
"dev": true,
"license": "Apache-2.0",
"bin": {
"playwright-core": "cli.js"
},
"engines": {
"node": ">=18"
}
},
"node_modules/typescript": {
"version": "5.9.3",
"resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz",
"integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==",
"dev": true,
"license": "Apache-2.0",
"bin": {
"tsc": "bin/tsc",
"tsserver": "bin/tsserver"
},
"engines": {
"node": ">=14.17"
}
},
"node_modules/undici-types": {
"version": "6.21.0",
"resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz",
"integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==",
"dev": true,
"license": "MIT"
}
}
}

16
tests/e2e/ui/package.json Normal file
View file

@ -0,0 +1,16 @@
{
"name": "litellm-ui-e2e",
"version": "0.0.0",
"private": true,
"scripts": {
"e2e": "playwright test --config playwright.config.ts",
"e2e:ui": "playwright test --ui --config playwright.config.ts",
"e2e:migration": "playwright test tests/migration/migratedPages.spec.ts --config playwright.config.ts",
"e2e:migration:root": "playwright test --config migration.serverRootPath.config.ts"
},
"devDependencies": {
"@playwright/test": "1.58.1",
"@types/node": "20.19.37",
"typescript": "5.9.3"
}
}

View file

@ -20,8 +20,8 @@ set -euo pipefail
# ================================================================
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
DASHBOARD_DIR="$(cd "$SCRIPT_DIR/.." && pwd)"
REPO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)"
DASHBOARD_DIR="$REPO_ROOT/ui/litellm-dashboard"
IS_CI="${CI:-false}"
CONTAINER_NAME="litellm-e2e-postgres-$$"
MOCK_PID=""
@ -187,12 +187,12 @@ PGPASSWORD="$DB_PASS" psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAM
# --- Playwright ---
echo "=== Installing Playwright dependencies ==="
cd "$DASHBOARD_DIR"
cd "$SCRIPT_DIR"
npm install --silent 2>/dev/null || true
npx playwright install chromium --with-deps 2>/dev/null || npx playwright install chromium
echo "=== Running Playwright tests ==="
npx playwright test --config e2e_tests/playwright.config.ts "$@"
npx playwright test --config playwright.config.ts "$@"
EXIT_CODE=$?
exit $EXIT_CODE

View file

@ -9,8 +9,9 @@ the default mount and a non-root `SERVER_ROOT_PATH` mount.
## Adding a page
When a page's migration merges, add its route segment to
`e2e_tests/fixtures/migratedPages.ts` (keep it in lockstep with `MIGRATED_PAGES`
in `src/utils/migratedPages.ts`). Both suites pick it up automatically.
`tests/e2e/ui/fixtures/migratedPages.ts` (keep it in lockstep with `MIGRATED_PAGES`
in `ui/litellm-dashboard/src/utils/migratedPages.ts`). Both suites pick it up
automatically.
## Running

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