Merge branch 'litellm_internal_staging' into litellm_mcp_server_env_vars

This commit is contained in:
Mateo Wang 2026-06-04 19:16:51 -07:00 • committed by GitHub
commit 201f40f608
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 1576 additions and 629 deletions

View file

@ -452,6 +452,120 @@ jobs:
- auth_ui_unit_tests_coverage.xml
- auth_ui_unit_tests_coverage
proxy_behavior_tests:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Seed DB schema via prisma db push
command: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Run proxy management behavior tests
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_behavior \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
proxy_security_tests:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Seed DB schema via prisma db push
command: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Run proxy security tests
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_security_tests \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
schema_migration_check:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
# An empty database; the test applies every committed migration itself.
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Check schema.prisma is in sync with committed migrations
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_migration_tests \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
litellm_router_testing: # Runs all tests with the "router" keyword
docker:
- *python312_image
@ -2643,6 +2757,12 @@ workflows:
filters: *main_branches
- auth_ui_unit_tests:
filters: *main_branches
- proxy_behavior_tests:
filters: *main_branches
- proxy_security_tests:
filters: *main_branches
- schema_migration_check:
filters: *main_branches
- build_docker_database_image:
filters: *main_branches
- e2e_ui_testing:

View file

@ -27,6 +27,11 @@ on:
required: false
type: number
default: 10
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: true
@ -82,18 +87,31 @@ jobs:
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
DIST: ${{ inputs.dist }}
run: |
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist=loadscope \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
else
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist="${DIST}" \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
fi
- name: Save coverage report
if: always()

View file

@ -1,190 +0,0 @@
name: _Unit Test Services Base (Reusable)
on:
workflow_call:
inputs:
test-path:
description: "Pytest path(s) to run"
required: true
type: string
workers:
description: "Number of pytest-xdist workers (0 = no parallelism)"
required: false
type: number
default: 2
reruns:
description: "Number of reruns for flaky tests"
required: false
type: number
default: 2
timeout-minutes:
description: "Job timeout in minutes"
required: false
type: number
default: 20
max-failures:
description: "Stop after this many failures"
required: false
type: number
default: 10
enable-postgres:
description: "Start a local Postgres service container and run Prisma migrations"
required: false
type: boolean
default: false
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: false
type: string
default: "run"
permissions:
contents: read
# The postgres service container below is spawned per-job on localhost and
# destroyed with the job. Nothing outside the runner can reach it. The
# user/password/database here are not secrets — they're bootstrap values
# for a throwaway container — so we hardcode them instead of attaching
# every matrix shard to a GHA environment just to read three "secrets"
# (which also produces a "temporarily deployed to …" notification on the
# PR timeline per shard per push).
jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.timeout-minutes }}
services:
postgres:
image: postgres@sha256:705a5d5b5836f3fcba0d02c4d281e6a7dd9ed2dd4078640f08a1e1e9896e097d # postgres:14
env:
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
POSTGRES_DB: litellm_test
ports:
- 5432:5432
options: >-
--health-cmd "pg_isready"
--health-interval 10s
--health-timeout 5s
--health-retries 5
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-services-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-services-
- name: Install dependencies
run: |
uv sync --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run Prisma migrations
if: ${{ inputs.enable-postgres }}
env:
DATABASE_URL: "postgresql://litellm:litellm@localhost:5432/litellm_test"
run: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- name: Run tests
env:
TEST_PATH: ${{ inputs.test-path }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
DIST: ${{ inputs.dist }}
DATABASE_URL: ${{ inputs.enable-postgres && 'postgresql://litellm:litellm@localhost:5432/litellm_test' || '' }}
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
else
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist="${DIST}" \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
fi
- name: Save coverage report
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage.xml
retention-days: 1
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always()
runs-on: ubuntu-latest
permissions:
contents: read
id-token: write
pull-requests: write
steps:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Download coverage report
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
with:
pattern: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage-reports
merge-multiple: true
- name: Upload to Codecov
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
with:
use_oidc: true
directory: coverage-reports
root_dir: ${{ github.workspace }}
flags: ${{ inputs.artifact-name }}
fail_ci_if_error: false

View file

@ -1,9 +1,10 @@
name: "Unit Tests: Proxy DB Operations"
# Uses DATABASE_URL secret — only runs on trusted branches, not PRs.
on:
push:
branches: [main, "litellm_**"]
pull_request:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
@ -30,9 +31,6 @@ concurrency:
# xdist balances its 188 parametrized cases across workers instead of
# pinning the whole file to one worker (the default --dist=loadscope
# behavior for single-file targets).
# * test_db_schema_migration.py is isolated because one test in it
# (test_aaaasschema_migration_check) takes ~170s — by itself it
# determines the shard's wall-clock floor.
jobs:
# Fast guard — fails the workflow if a test_*.py file under
# tests/proxy_unit_tests/ is not referenced by any matrix entry below.
@ -166,18 +164,6 @@ jobs:
dist: loadscope
timeout: 15
# ---- db-and-spend: isolate the 170s schema-migration test ----
# test_db_schema_migration.py has exactly one test, and that test
# is mostly waiting on `prisma migrate deploy` / `prisma migrate
# diff` subprocesses (~170s). It does no CPU-bound Python work
# inside the test. Running with workers=0 (serial, no xdist)
# skips the 4-worker cold-start cost we'd otherwise pay for a
# single test, saving ~4 minutes of wall-clock.
- test-group: schema-migration
test-path: "tests/proxy_unit_tests/test_db_schema_migration.py"
workers: 0
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: >-
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
@ -232,12 +218,11 @@ jobs:
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-services-base.yml
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
workers: ${{ matrix.workers }}
reruns: 2
timeout-minutes: ${{ matrix.timeout }}
enable-postgres: true
dist: ${{ matrix.dist }}
artifact-name: proxy-db-${{ matrix.test-group }}

View file

@ -1,34 +0,0 @@
name: "Unit Tests: Proxy Management-Endpoint Behavior Pinning"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
proxy-mgmt-behavior:
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: tests/proxy_behavior
# workers=0 (no xdist): the world seed is a single shared Postgres
# state — two xdist workers both call seed_world() and race on the
# ``behavior-pin-budget`` row, producing UniqueViolation + cascading
# missing-membership FK failures. The whole suite is ~7s sequentially,
# so the cost of disabling parallelism here is negligible.
workers: 0
reruns: 0
enable-postgres: true
artifact-name: proxy-mgmt-behavior
timeout-minutes: 15

View file

@ -1,28 +0,0 @@
name: "Unit Tests: Security"
# Kept push-only (was previously required by DATABASE_URL secret scoping;
# now the postgres credentials are ephemeral localhost values but the
# push-trigger stays to match the proxy-db workflow cadence).
on:
push:
branches: [main, "litellm_**"]
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
security:
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: "tests/proxy_security_tests/"
workers: 1
reruns: 2
timeout-minutes: 20
enable-postgres: true
artifact-name: security

View file

@ -4508,6 +4508,25 @@ class JWTRoutingOverride(BaseModel):
}
class UnregisteredJWTClientBehavior(str, enum.Enum):
"""
Controls what happens when `virtual_key_claim_field` is configured but the
JWT claim value has no registered mapping in `litellm_jwtkeymapping`.
- fallback_team_mapping: Fall through to standard team-based JWT auth (default,
backward-compatible).
- reject: Immediately return HTTP 403. Use this when every valid JWT client
must have a pre-registered virtual key — unknown callers are denied.
- auto_register: Automatically create a new virtual key and mapping on first
encounter. The new key has no budget/model restrictions; admins can tighten
it later via /jwt_client/update.
"""
FALLBACK_TEAM_MAPPING = "fallback_team_mapping"
REJECT = "reject"
AUTO_REGISTER = "auto_register"
class JWTIssuerConfig(BaseModel):
"""
Issuer-bound JWT validation configuration.
@ -4674,6 +4693,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
default=300,
description="TTL (seconds) for caching JWT-to-virtual-key mapping lookups.",
)
unregistered_jwt_client_behavior: UnregisteredJWTClientBehavior = Field(
default=UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING,
description=(
"What to do when virtual_key_claim_field is set but the JWT claim value "
"has no registered mapping. 'fallback_team_mapping' (default): fall through "
"to team-based JWT auth. 'reject': return HTTP 403. "
"'auto_register': auto-create a virtual key and mapping on first encounter."
),
)
routing_overrides: Optional[List[JWTRoutingOverride]] = Field(
default=None,
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
@ -4693,6 +4721,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
# ``s3://`` / ``gcs://`` when this is None.
config_file_path = kwargs.pop("config_file_path", None)
# Backward-compat: jwt_client_id_field was renamed to virtual_key_claim_field
if "jwt_client_id_field" in kwargs:
if "virtual_key_claim_field" not in kwargs:
kwargs["virtual_key_claim_field"] = kwargs.pop("jwt_client_id_field")
else:
kwargs.pop("jwt_client_id_field")
# get the attribute names for this Pydantic model
allowed_keys = LiteLLM_JWTAuth.__annotations__.keys()

View file

@ -12,7 +12,7 @@ import fnmatch
import re
import secrets
from datetime import datetime, timezone
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union, cast
from typing import Any, Dict, Iterator, NamedTuple, List, Optional, Tuple, Union, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
@ -607,6 +607,169 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
return api_key
# Cache sentinel written when a JWT under AUTO_REGISTER resolved to a proxy
# admin via auth_builder. Proxy admins don't need a mapped virtual key (they
# have full access via auth_builder anyway), but without a cache entry every
# subsequent request from the same JWT identity would re-query the DB for a
# non-existent mapping. Sentinel tells _resolve_jwt_to_virtual_key to skip
# the lookup and return None (caller proceeds to auth_builder).
_JWT_PROXY_ADMIN_SENTINEL = "__JWT_PROXY_ADMIN__"
class _PendingAutoRegister(NamedTuple):
"""
Signal returned by ``_resolve_jwt_to_virtual_key`` when the JWT's claim is
unmapped and ``unregistered_jwt_client_behavior`` is AUTO_REGISTER.
The caller MUST run standard ``JWTAuthManager.auth_builder`` to apply RBAC,
scope mappings, ``custom_validate``, and ``user_allowed_email_domain``
policy BEFORE calling ``_auto_register_jwt_mapping`` with the validated
``team_id`` / ``user_id`` from the auth_builder result. Auto-registering
purely on a signature-valid JWT (the old behavior) bypassed every JWT
policy beyond signature verification.
"""
claim_field: str
claim_value: str
cache_key: str
async def _auto_register_jwt_mapping(
virtual_key_claim_field: str,
claim_value: str,
jwt_handler: JWTHandler,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
cache_key: str,
team_id: Optional[str] = None,
user_id: Optional[str] = None,
org_id: Optional[str] = None,
end_user_id: Optional[str] = None,
) -> Optional[UserAPIKeyAuth]:
"""
Auto-register: create a new virtual key + mapping for an unrecognised JWT
claim value. ``team_id`` and ``user_id`` must come from a successful
``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER
RBAC/scope/custom_validate/email-domain policy has been enforced. The key
is stamped with those values so the cached future-request path inherits
the same team/user/org limits the auth_builder path would have applied.
Race safety: if two concurrent requests both reach here simultaneously (both
saw no mapping in the DB), one will win the unique-constraint race on
litellm_jwtkeymapping. The loser catches the conflict, deletes its orphaned
key, fetches the winner's mapping, and proceeds — no error surfaced.
"""
# Inline import required: key_management_endpoints imports user_api_key_auth
# (line 51) so a module-level import here would create a circular dependency.
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
)
# ``table_name="key"`` is required: without it, generate_key_helper_fn
# falls into the user-upsert branch (`table_name is None or "user"`) and
# attempts to insert into LiteLLM_UserTable with user_id=None, which fails
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
# /key/generate) passes table_name="key" explicitly.
key_data = await generate_key_helper_fn(
request_type="key",
table_name="key",
team_id=team_id,
user_id=user_id,
organization_id=org_id,
metadata={
"auto_registered": True,
"jwt_claim_field": virtual_key_claim_field,
"jwt_claim_value": claim_value,
},
)
# generate_key_helper_fn returns the plaintext key in "token"; the persisted
# row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK
# value referenced by LiteLLM_JWTKeyMapping.token.
token_hash = hash_token(key_data["token"])
try:
await prisma_client.db.litellm_jwtkeymapping.create(
data={
"jwt_claim_name": virtual_key_claim_field,
"jwt_claim_value": claim_value,
"token": token_hash,
"created_by": "auto_register",
"updated_by": "auto_register",
}
)
except Exception as e:
error_str = str(e).lower()
if "unique" in error_str or "p2002" in error_str:
# A concurrent request won the race. The key generate_key_helper_fn
# just persisted to LiteLLM_VerificationToken is orphaned — nothing
# maps to it, but it's a fully valid unrestricted API key sitting in
# the DB and the cleartext is in memory on this request. Delete it
# so orphans don't accumulate under sustained concurrency.
verbose_proxy_logger.debug(
"JWT Key Mapping (auto_register): unique conflict on create — "
"deleting orphaned virtual key and fetching winner's mapping for %s='%s'.",
virtual_key_claim_field,
claim_value,
)
try:
await prisma_client.db.litellm_verificationtoken.delete(
where={"token": token_hash}
)
except Exception as delete_err:
# Don't fail the request if cleanup fails — the orphan is
# unmapped and inert. Log so an operator can prune it later.
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s",
delete_err,
)
token_hash = await get_jwt_key_mapping_object(
jwt_claim_name=virtual_key_claim_field,
jwt_claim_value=claim_value,
prisma_client=prisma_client,
)
if token_hash is None:
# The winner's mapping vanished between the unique-constraint
# conflict and our re-fetch (concurrent delete). Returning None
# here would silently fall through to team-based JWT auth —
# a less-restrictive path than the operator configured. Raise
# 503 so the caller retries against a stable state instead.
raise HTTPException(
status_code=503,
detail=(
"JWT Key Mapping: AUTO_REGISTER race resolution failed — "
"winner's mapping was concurrently removed. Retry the request."
),
)
else:
raise
await user_api_key_cache.async_set_cache(
key=cache_key,
value=token_hash,
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
verbose_proxy_logger.info(
"JWT Key Mapping (auto_register): created new virtual key for %s='%s'.",
virtual_key_claim_field,
claim_value,
)
auto_registered_key = await get_key_object(
hashed_token=token_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if auto_registered_key is not None:
auto_registered_key.org_id = org_id
auto_registered_key.end_user_id = end_user_id
return auto_registered_key
async def _resolve_jwt_to_virtual_key(
jwt_claims: dict,
jwt_handler: JWTHandler,
@ -614,7 +777,22 @@ async def _resolve_jwt_to_virtual_key(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
) -> Optional[UserAPIKeyAuth]:
) -> Union[Optional[UserAPIKeyAuth], "_PendingAutoRegister"]:
"""
Returns:
- ``UserAPIKeyAuth``: a resolved virtual key (cache hit or DB hit). The
caller may use this directly; JWT policy has been enforced previously
(at key-creation time or, for cached results, before caching).
- ``_PendingAutoRegister``: claim is unmapped and behavior is AUTO_REGISTER.
The caller MUST run ``JWTAuthManager.auth_builder`` to enforce JWT
policy (RBAC, scope, custom_validate, email-domain), then invoke
``_auto_register_jwt_mapping`` with the validated team_id/user_id.
- ``None``: claim is unmapped and behavior is FALLBACK_TEAM_MAPPING.
The caller falls through to standard team-based JWT auth (which itself
enforces full JWT policy via auth_builder).
- Raises HTTPException: REJECT policy hit, missing claim under
REJECT/AUTO_REGISTER, or other policy violations.
"""
virtual_key_claim_field = jwt_handler.litellm_jwtauth.virtual_key_claim_field
if virtual_key_claim_field is None:
return None
@ -629,12 +807,61 @@ async def _resolve_jwt_to_virtual_key(
verbose_proxy_logger.debug(
f"JWT Key Mapping: Claim field '{virtual_key_claim_field}' not found in JWT claims."
)
# A missing claim is an unmapped client — apply the no-match policy
# rather than returning early. Otherwise a caller can bypass REJECT
# simply by presenting a JWT that omits the configured field. For
# AUTO_REGISTER there is no stable identity to map without a claim
# value, so we deny rather than create a sentinel-keyed record.
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
if behavior in (
UnregisteredJWTClientBehavior.REJECT,
UnregisteredJWTClientBehavior.AUTO_REGISTER,
):
raise HTTPException(
status_code=403,
detail=(
f"JWT Key Mapping: Required claim '{virtual_key_claim_field}' "
"is missing from the JWT. Access denied."
),
)
return None
cache_key = f"jwt_key_mapping:{virtual_key_claim_field}:{claim_value}"
cached_mapping = await user_api_key_cache.async_get_cache(cache_key)
if cached_mapping == _JWT_PROXY_ADMIN_SENTINEL:
# Previously resolved to a proxy admin via auth_builder; skip the
# mapping lookup and let the caller re-run auth_builder. Avoids a
# repeated DB hit on every proxy-admin request under AUTO_REGISTER.
return None
if cached_mapping == "__NO_MAPPING__":
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
if behavior == UnregisteredJWTClientBehavior.REJECT:
raise HTTPException(
status_code=403,
detail=f"JWT Key Mapping: No registered mapping for {virtual_key_claim_field}='{claim_value}'. Access denied.",
)
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER:
# Stale sentinel written under a prior fallback_team_mapping config —
# evict it and defer auto-register to after auth_builder runs. Raise
# the same 500 as the fresh-path AUTO_REGISTER branch when there is
# no DB, so behavior is consistent regardless of whether the cache
# happens to hold the sentinel.
if prisma_client is None:
raise HTTPException(
status_code=500,
detail=(
"JWT Key Mapping: AUTO_REGISTER requires a database connection. "
"Configure a database or change unregistered_jwt_client_behavior."
),
)
await user_api_key_cache.async_delete_cache(cache_key)
return _PendingAutoRegister(
claim_field=virtual_key_claim_field,
claim_value=str(claim_value),
cache_key=cache_key,
)
return None
elif cached_mapping is not None:
return await get_key_object(
@ -645,14 +872,15 @@ async def _resolve_jwt_to_virtual_key(
proxy_logging_obj=proxy_logging_obj,
)
if prisma_client is None:
return None
token_hash = await get_jwt_key_mapping_object(
jwt_claim_name=virtual_key_claim_field,
jwt_claim_value=str(claim_value),
prisma_client=prisma_client,
)
# Resolve the mapping from DB, or treat prisma_client=None as a definitive
# miss (no DB → no mapping can exist → apply no-match policy below).
token_hash: Optional[str] = None
if prisma_client is not None:
token_hash = await get_jwt_key_mapping_object(
jwt_claim_name=virtual_key_claim_field,
jwt_claim_value=str(claim_value),
prisma_client=prisma_client,
)
if token_hash is not None:
await user_api_key_cache.async_set_cache(
@ -667,13 +895,50 @@ async def _resolve_jwt_to_virtual_key(
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
else:
# No mapping found (DB miss or no DB) — apply no-match policy.
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
if behavior == UnregisteredJWTClientBehavior.REJECT:
# Cache the miss before raising so repeated rejections are served from
# cache and don't re-query the DB on every request.
await user_api_key_cache.async_set_cache(
key=cache_key,
value="__NO_MAPPING__",
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return None
raise HTTPException(
status_code=403,
detail=f"JWT Key Mapping: No registered mapping for {virtual_key_claim_field}='{claim_value}'. Access denied.",
)
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER:
if prisma_client is None:
raise HTTPException(
status_code=500,
detail=(
"JWT Key Mapping: AUTO_REGISTER requires a database connection. "
"Configure a database or change unregistered_jwt_client_behavior."
),
)
# Defer: caller runs JWTAuthManager.auth_builder to enforce RBAC, scope,
# custom_validate, and email-domain policy, then auto-registers using
# the validated identity. Auto-registering here on a signature-only
# JWT would bypass every JWT policy beyond signature verification.
return _PendingAutoRegister(
claim_field=virtual_key_claim_field,
claim_value=str(claim_value),
cache_key=cache_key,
)
# FALLBACK_TEAM_MAPPING (default): cache the miss and return None so the
# caller falls through to standard team-based JWT auth.
await user_api_key_cache.async_set_cache(
key=cache_key,
value="__NO_MAPPING__",
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return None
def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
@ -893,6 +1158,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Try JWT-to-Virtual-Key mapping first to avoid
# unnecessary DB queries in auth_builder
do_standard_jwt_auth = True
pending_auto_register: Optional[_PendingAutoRegister] = None
if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None:
# Decode JWT to get claims without running full auth_builder
jwt_claims: Optional[dict]
@ -901,7 +1167,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
else:
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
valid_token = await _resolve_jwt_to_virtual_key(
resolve_result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
@ -909,11 +1175,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if valid_token is not None:
if isinstance(resolve_result, UserAPIKeyAuth):
valid_token = resolve_result
api_key = valid_token.token or ""
valid_token.jwt_claims = jwt_claims
do_standard_jwt_auth = False
# Fall through to virtual key checks
elif isinstance(resolve_result, _PendingAutoRegister):
# Run full JWT policy (RBAC, scope, custom_validate,
# email-domain) via auth_builder, then create the key
# from the validated identity below.
pending_auto_register = resolve_result
# else: None → FALLBACK_TEAM_MAPPING, falls through to
# standard JWT auth_builder below
if do_standard_jwt_auth:
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
@ -946,6 +1220,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
jwt_claims = result.get("jwt_claims", None)
if is_proxy_admin:
# Proxy admins authenticate via auth_builder (full
# access), not via a mapped virtual key. If
# AUTO_REGISTER was pending, cache a sentinel so
# future requests from this JWT identity skip the
# DB mapping lookup in _resolve_jwt_to_virtual_key.
# Without this, every proxy-admin request under
# AUTO_REGISTER re-hits get_jwt_key_mapping_object.
if pending_auto_register is not None:
await user_api_key_cache.async_set_cache(
key=pending_auto_register.cache_key,
value=_JWT_PROXY_ADMIN_SENTINEL,
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return UserAPIKeyAuth(
api_key=None,
user_role=LitellmUserRoles.PROXY_ADMIN,
@ -1032,6 +1319,32 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
else None
)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
# JWT policy (RBAC, scope, custom_validate, email-domain)
# has now been enforced by auth_builder above. Create the
# mapping + virtual key from the *validated* identity, then
# replace valid_token with the new key so downstream checks
# use the key-scoped path.
if pending_auto_register is not None and prisma_client is not None:
auto_registered = await _auto_register_jwt_mapping(
virtual_key_claim_field=pending_auto_register.claim_field,
claim_value=pending_auto_register.claim_value,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
cache_key=pending_auto_register.cache_key,
team_id=team_id,
user_id=user_id,
org_id=org_id,
end_user_id=end_user_id,
)
if auto_registered is not None:
auto_registered.jwt_claims = jwt_claims
valid_token = auto_registered
api_key = valid_token.token or ""
# Check if model has zero cost - if so, skip all budget checks
model = _get_model_from_request_context(
request_data=request_data,

View file

@ -0,0 +1,70 @@
import os
import shutil
import subprocess
import tempfile
from pathlib import Path
import pytest
@pytest.mark.skipif(
"DATABASE_URL" not in os.environ,
reason="requires a postgres database (DATABASE_URL)",
)
def test_schema_migration_in_sync():
"""Fail if schema.prisma has changes not captured by the committed migrations.
Applies every committed migration to an empty database, then diffs the result
against schema.prisma. A non-empty diff means the schema was changed without a
matching migration being generated.
"""
db_url = os.environ["DATABASE_URL"]
source_migrations_dir = Path(
"./litellm-proxy-extras/litellm_proxy_extras/migrations"
)
source_schema_path = Path("./schema.prisma")
temp_base = Path(tempfile.mkdtemp(prefix="litellm_schema_migration_"))
schema_path = temp_base / "schema.prisma"
migrations_dir = temp_base / "migrations"
try:
shutil.copy(source_schema_path, schema_path)
shutil.copytree(source_migrations_dir, migrations_dir)
if not any(migrations_dir.iterdir()):
pytest.fail(
"No existing migrations found. Run `python litellm/ci_cd/baseline_db_migration.py`."
)
subprocess.run(
["prisma", "migrate", "deploy", "--schema", str(schema_path)],
check=True,
env={**os.environ, "DATABASE_URL": db_url},
)
diff = subprocess.run(
[
"prisma",
"migrate",
"diff",
"--from-url",
db_url,
"--to-schema-datamodel",
str(schema_path),
"--script",
"--exit-code",
],
capture_output=True,
text=True,
)
if diff.returncode == 2:
pytest.fail(
"Schema changes detected that no migration captures. Run "
"`python litellm/ci_cd/run_migration.py <migration_name>`.\n\n"
+ diff.stdout
)
assert diff.returncode == 0, f"prisma migrate diff errored: {diff.stderr}"
finally:
shutil.rmtree(temp_base, ignore_errors=True)

View file

@ -1,39 +1,32 @@
import os
import pytest
from fastapi.testclient import TestClient
from litellm.proxy.proxy_server import app, ProxyLogging
from litellm.proxy.proxy_server import app, ProxyLogging, hash_token
from litellm.caching import DualCache
MASTER_KEY = "sk-1234"
@pytest.fixture(autouse=True)
def override_env_settings(monkeypatch):
# Set environment variables only for tests using-monkeypatch (function scope by default).
# Use DATABASE_URL from environment (set by CircleCI to local postgres)
if "DATABASE_URL" not in os.environ:
pytest.fail(
"DATABASE_URL not set - this test requires a local postgres database to be running"
"DATABASE_URL not set - this test requires a postgres database to be running"
)
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
monkeypatch.setenv("LITELLM_MASTER_KEY", MASTER_KEY)
monkeypatch.setenv("LITELLM_LOG", "DEBUG")
@pytest.fixture(scope="module")
def test_client():
"""
This fixture starts up the test client which triggers FastAPI's startup events.
Prisma will connect to the DB using the provided DATABASE_URL.
"""
"""Starting the test client triggers FastAPI startup, where Prisma connects to the DB."""
with TestClient(app) as client:
yield client
@pytest.mark.asyncio
async def test_master_key_not_inserted(test_client):
"""
This test ensures that when the app starts (or when you hit the /health endpoint
to trigger startup logic), no unexpected write occurs in the DB.
"""
# Hit an endpoint (like /health) that triggers any startup tasks.
"""The master key must never be persisted to the verification-token table on startup."""
response = test_client.get("/health/liveliness")
assert response.status_code == 200
@ -46,13 +39,22 @@ async def test_master_key_not_inserted(test_client):
),
)
# Connect directly to the test database to inspect the data.
await prisma_client.connect()
result = await prisma_client.db.litellm_verificationtoken.find_many()
print(result)
stored_tokens = {
row.token
for row in await prisma_client.db.litellm_verificationtoken.find_many()
}
# The expectation is that no token (or unintended record) is added on startup.
assert len(result) == 0, (
"SECURITY ALERT SECURITY ALERT SECURITY ALERT: Expected no record in the litellm_verificationtoken table. On startup - the master key should NOT be Inserted into the DB."
"We have found keys in the DB. This is unexpected and should not happen."
)
for leaked in (hash_token(MASTER_KEY), MASTER_KEY):
assert leaked not in stored_tokens, (
"SECURITY ALERT: the master key was found in the litellm_verificationtoken "
"table. The master key must never be inserted into the DB."
)
# Canary against any other unexpected startup write (default key, rotation
# artifact, ...). The job gives each run a fresh DB, so a clean startup must
# leave the table empty; if startup ever legitimately seeds a token, narrow
# this while keeping the master-key assertion above.
assert (
not stored_tokens
), f"startup unexpectedly wrote token(s) to litellm_verificationtoken: {stored_tokens}"

View file

@ -1,87 +0,0 @@
import pytest
import os
import subprocess
from pathlib import Path
from pytest_postgresql import factories
import shutil
import tempfile
# Create postgresql fixture
postgresql_my_proc = factories.postgresql_proc(port=None)
postgresql_my = factories.postgresql("postgresql_my_proc")
@pytest.fixture(scope="function")
def schema_setup(postgresql_my):
"""Fixture to provide a test postgres database"""
return postgresql_my
@pytest.mark.xdist_group("proxy_heavy")
def test_aaaasschema_migration_check(schema_setup, monkeypatch):
"""Test to check if schema requires migration"""
# Set test database URL
test_db_url = f"postgresql://{schema_setup.info.user}:@{schema_setup.info.host}:{schema_setup.info.port}/{schema_setup.info.dbname}"
# test_db_url = "postgresql://test-user:test-password@test-host.example.com/test-db?sslmode=require"
monkeypatch.setenv("DATABASE_URL", test_db_url)
deploy_dir = Path("./litellm-proxy-extras/litellm_proxy_extras")
source_migrations_dir = deploy_dir / "migrations"
source_schema_path = Path("./schema.prisma")
# Use worker-specific temp directory to avoid races when running with -n 8.
# Prisma expects migrations in <schema_dir>/migrations, so we create that layout.
temp_base = Path(tempfile.mkdtemp(prefix="litellm_schema_migration_"))
temp_migrations_dir = temp_base / "migrations"
schema_path = temp_base / "schema.prisma"
try:
shutil.copy(source_schema_path, schema_path)
shutil.copytree(source_migrations_dir, temp_migrations_dir)
if not temp_migrations_dir.exists() or not any(temp_migrations_dir.iterdir()):
print("No existing migrations found - first migration needed")
pytest.fail(
"No existing migrations found - first migration needed. Run `litellm/ci_cd/baseline_db.py` to create new migration -E.g. `python litellm/ci_cd/baseline_db_migration.py`."
)
# Apply all existing migrations
subprocess.run(
["prisma", "migrate", "deploy", "--schema", str(schema_path)], check=True
)
# Compare current database state against schema
diff_result = subprocess.run(
[
"prisma",
"migrate",
"diff",
"--from-url",
test_db_url,
"--to-schema-datamodel",
str(schema_path),
"--script", # Show the SQL diff
"--exit-code", # Return exit code 2 if there are differences
],
capture_output=True,
text=True,
)
print("Exit code:", diff_result.returncode)
print("Stdout:", diff_result.stdout)
print("Stderr:", diff_result.stderr)
if diff_result.returncode == 2:
print("Schema changes detected. New migration needed.")
print("Schema differences:")
print(diff_result.stdout)
pytest.fail(
"Schema changes detected - new migration required. Run `litellm/ci_cd/run_migration.py` to create new migration -E.g. `python litellm/ci_cd/run_migration.py <migration_name>`."
)
else:
print("No schema changes detected. Migration not needed.")
finally:
# Clean up: remove temporary directory
if temp_base.exists():
shutil.rmtree(temp_base)

View file

@ -747,7 +747,6 @@ async def test_allowed_routes_admin(
from litellm.proxy.proxy_server import user_api_key_auth
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
await litellm.proxy.proxy_server.prisma_client.connect()
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://example.com/public-key")

View file

@ -27,7 +27,6 @@ from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
from litellm.caching.caching import DualCache
from fastapi import HTTPException
# ──────────────────────────────────────────────
# Tests: _resolve_jwt_to_virtual_key
# ──────────────────────────────────────────────
@ -454,3 +453,856 @@ async def test_create_success_returns_response_without_token():
assert isinstance(result, JWTKeyMappingResponse)
assert "token" not in result.model_fields
assert result.jwt_claim_name == "email"
# ──────────────────────────────────────────────
# Tests: unregistered_jwt_client_behavior
# ──────────────────────────────────────────────
@pytest.mark.asyncio
async def test_reject_behavior_raises_403_on_no_mapping():
"""
When unregistered_jwt_client_behavior='reject' and no mapping exists,
_resolve_jwt_to_virtual_key must raise HTTP 403.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.REJECT,
)
jwt_claims = {"email": "unknown@example.com"}
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
user_api_key_cache = DualCache()
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
assert "unknown@example.com" in exc_info.value.detail
@pytest.mark.asyncio
async def test_reject_behavior_caches_sentinel_after_db_miss():
"""
On a fresh DB miss with REJECT, the __NO_MAPPING__ sentinel must be written
to cache so that subsequent rejected requests are served from cache and do
not re-query the DB.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.REJECT,
virtual_key_mapping_cache_ttl=300,
)
jwt_claims = {"email": "unknown@example.com"}
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
user_api_key_cache = DualCache()
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
# First call — DB miss, should raise 403 and write sentinel
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
# Sentinel must now be in cache
cached = await user_api_key_cache.async_get_cache(
"jwt_key_mapping:email:unknown@example.com"
)
assert cached == "__NO_MAPPING__"
# Second call — must raise 403 from cache, no additional DB hit
prisma_client.db.litellm_jwtkeymapping.find_first.reset_mock()
with pytest.raises(HTTPException) as exc_info2:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info2.value.status_code == 403
prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called()
@pytest.mark.asyncio
async def test_reject_behavior_raises_403_on_cached_no_mapping():
"""
When the negative-cache sentinel __NO_MAPPING__ is present and behavior is
'reject', the function must also raise HTTP 403 (not return None silently).
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.REJECT,
)
jwt_claims = {"email": "unknown@example.com"}
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
# Pre-populate the negative cache so the DB is not hit
user_api_key_cache = DualCache()
cache_key = "jwt_key_mapping:email:unknown@example.com"
await user_api_key_cache.async_set_cache(cache_key, "__NO_MAPPING__")
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
# DB must NOT have been hit (sentinel served from cache)
prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called()
@pytest.mark.asyncio
async def test_auto_register_returns_pending_signal_without_creating_key():
"""
Security: when unregistered_jwt_client_behavior='auto_register' and no
mapping exists, _resolve_jwt_to_virtual_key must NOT create the key yet.
It returns a _PendingAutoRegister signal so the caller can run
JWTAuthManager.auth_builder (enforcing RBAC, scope mappings,
custom_validate, user_allowed_email_domain) FIRST. Creating the key here
would bypass every JWT policy beyond signature verification.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
jwt_claims = {"sub": "new-user-42"}
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = DualCache()
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_gen_key:
result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert isinstance(result, _PendingAutoRegister)
assert result.claim_field == "sub"
assert result.claim_value == "new-user-42"
assert result.cache_key == "jwt_key_mapping:sub:new-user-42"
# CRITICAL: no key was created — that must wait until after auth_builder
mock_gen_key.assert_not_called()
prisma_client.db.litellm_jwtkeymapping.create.assert_not_called()
@pytest.mark.asyncio
async def test_auto_register_creates_key_and_mapping_when_helper_invoked():
"""
When the caller invokes _auto_register_jwt_mapping directly (after
auth_builder validation), the helper creates the key + mapping row and
returns a UserAPIKeyAuth. The mapping row stores the hashed token (FK to
LiteLLM_VerificationToken), not the plaintext key.
"""
from litellm.proxy._types import hash_token
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
virtual_key_mapping_cache_ttl=300,
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = DualCache()
plaintext_key = "sk-auto-key"
expected_hash = hash_token(plaintext_key)
mock_key_obj = UserAPIKeyAuth(token=expected_hash, team_id="validated-team")
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
) as mock_get_key,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_gen_key,
):
mock_gen_key.return_value = {"token": plaintext_key, "key": plaintext_key}
mock_get_key.return_value = mock_key_obj
result = await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="new-user-42",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
cache_key="jwt_key_mapping:sub:new-user-42",
team_id="validated-team",
user_id="validated-user",
)
assert result == mock_key_obj
# generate_key_helper_fn was passed table_name="key" (not user-upsert path)
# and the validated team_id + user_id from auth_builder
assert mock_gen_key.call_args.kwargs["table_name"] == "key"
assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team"
assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user"
# Mapping row was created with the hashed token (FK target)
call_data = prisma_client.db.litellm_jwtkeymapping.create.call_args[1]["data"]
assert call_data["jwt_claim_name"] == "sub"
assert call_data["jwt_claim_value"] == "new-user-42"
assert call_data["token"] == expected_hash
cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:new-user-42")
assert cached == expected_hash
@pytest.mark.asyncio
async def test_auto_register_returns_pending_signal_on_stale_no_mapping_sentinel():
"""
If the cache holds a stale __NO_MAPPING__ sentinel (written under a prior
fallback_team_mapping config) and behavior is now AUTO_REGISTER, the
resolver must evict the sentinel and return _PendingAutoRegister (so the
caller can run auth_builder before creating the key) — not silently return
None and not create the key on the spot.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
jwt_claims = {"email": "alice@corp.com"}
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = DualCache()
await user_api_key_cache.async_set_cache(
"jwt_key_mapping:email:alice@corp.com", "__NO_MAPPING__"
)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_gen_key:
result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert isinstance(result, _PendingAutoRegister)
# Stale sentinel must be evicted so the deferred auto-register actually
# runs after auth_builder validates the JWT
cached_after = await user_api_key_cache.async_get_cache(
"jwt_key_mapping:email:alice@corp.com"
)
assert cached_after is None
mock_gen_key.assert_not_called()
prisma_client.db.litellm_jwtkeymapping.create.assert_not_called()
@pytest.mark.asyncio
async def test_auto_register_race_condition_unique_conflict():
"""
If two concurrent requests both call _auto_register_jwt_mapping and the
second hits a unique-constraint violation on create, it must:
1) delete the orphaned virtual key it just created (so orphans don't
accumulate in LiteLLM_VerificationToken under sustained concurrency),
2) fall back to the winner's mapping,
3) not surface an error.
"""
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy._types import UnregisteredJWTClientBehavior, hash_token
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(
side_effect=Exception("Unique constraint failed (P2002)")
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock()
# Simulate the winner's mapping already in DB after the conflict
winner_mapping = MagicMock()
winner_mapping.token = "winner_token_hash"
winner_mapping.is_active = True
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(
return_value=winner_mapping
)
user_api_key_cache = DualCache()
loser_plaintext = "sk-loser"
loser_hash = hash_token(loser_plaintext)
mock_key_obj = UserAPIKeyAuth(token="winner_token_hash", team_id=None)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
) as mock_get_key,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": loser_plaintext, "key": loser_plaintext},
),
):
mock_get_key.return_value = mock_key_obj
result = await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="user-42",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
cache_key="jwt_key_mapping:sub:user-42",
)
assert result == mock_key_obj
# The orphaned loser key must be deleted from LiteLLM_VerificationToken
prisma_client.db.litellm_verificationtoken.delete.assert_called_once_with(
where={"token": loser_hash}
)
# Cache should hold the winner's token, not the loser's
cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:user-42")
assert cached == "winner_token_hash"
mock_get_key.assert_called_once_with(
hashed_token="winner_token_hash",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
# ──────────────────────────────────────────────
# Tests: prisma_client=None does not bypass no-match policy
# ──────────────────────────────────────────────
@pytest.mark.asyncio
async def test_reject_behavior_enforced_when_prisma_client_is_none():
"""
When prisma_client is None and behavior is REJECT, a 403 must be raised —
not silently fallen through to team auth.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.REJECT,
)
jwt_claims = {"email": "unknown@example.com"}
user_api_key_cache = DualCache()
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=None, # no DB
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
assert "unknown@example.com" in exc_info.value.detail
@pytest.mark.asyncio
async def test_reject_raises_403_when_claim_field_missing_from_jwt():
"""
Security: a JWT that omits the configured virtual_key_claim_field must NOT
bypass the REJECT policy. Previously the early `if claim_value is None:
return None` branch ran before the policy check, letting a caller who knows
the configured claim-field name silently fall through to team-based auth.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.REJECT,
)
# JWT does NOT contain "sub"
jwt_claims = {"email": "user@example.com"}
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
assert "'sub'" in exc_info.value.detail
assert "missing from the JWT" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auto_register_raises_403_when_claim_field_missing_from_jwt():
"""
AUTO_REGISTER cannot create a mapping without a stable identity. When the
configured claim field is missing from the JWT, return 403 rather than
silently falling through (which would bypass the unregistered-client policy)
or creating a sentinel-keyed record.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
)
jwt_claims = {"email": "user@example.com"}
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 403
assert "missing from the JWT" in exc_info.value.detail
@pytest.mark.asyncio
async def test_fallback_team_mapping_returns_none_when_claim_field_missing_from_jwt():
"""
Under FALLBACK_TEAM_MAPPING (the default, backward-compatible mode), a JWT
without the configured claim field must still fall through to team-based
JWT auth — not raise. This preserves the pre-existing contract.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING,
)
jwt_claims = {"email": "user@example.com"}
result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
)
assert result is None
@pytest.mark.asyncio
async def test_fallback_team_mapping_returns_none_when_prisma_client_is_none():
"""
When prisma_client is None and behavior is FALLBACK_TEAM_MAPPING, the
function must return None (fall through to team auth) — not raise.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="email",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING,
)
jwt_claims = {"email": "anyone@example.com"}
result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
)
assert result is None
@pytest.mark.asyncio
async def test_auto_register_raises_500_when_prisma_client_is_none():
"""
AUTO_REGISTER without a DB connection must raise HTTP 500 with a clear
message — it cannot create keys without a database.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
)
jwt_claims = {"sub": "new-user-42"}
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 500
assert "AUTO_REGISTER requires a database" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auto_register_raises_500_when_sentinel_cached_and_no_db():
"""
AUTO_REGISTER + cached __NO_MAPPING__ sentinel + prisma_client is None must
raise HTTP 500, matching the fresh-path behavior. Previously this path
silently returned None and let the request fall through to team auth,
creating different access-control outcomes under identical configuration.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
jwt_claims = {"sub": "user-42"}
user_api_key_cache = DualCache()
# Stale sentinel written under a prior fallback_team_mapping config
await user_api_key_cache.async_set_cache(
"jwt_key_mapping:sub:user-42", "__NO_MAPPING__"
)
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 500
assert "AUTO_REGISTER requires a database" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auto_register_race_conflict_tolerates_delete_failure():
"""
If deleting the orphaned virtual key after a race-condition conflict fails
(e.g. transient DB error), the request must still succeed by returning the
winner's mapping — the orphan is unmapped and inert.
"""
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(
side_effect=Exception("Unique constraint failed (P2002)")
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock(
side_effect=Exception("transient DB error")
)
winner_mapping = MagicMock()
winner_mapping.token = "winner_token_hash"
winner_mapping.is_active = True
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(
return_value=winner_mapping
)
user_api_key_cache = DualCache()
mock_key_obj = UserAPIKeyAuth(token="winner_token_hash", team_id=None)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
) as mock_get_key,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": "sk-loser", "key": "sk-loser"},
),
):
mock_get_key.return_value = mock_key_obj
result = await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="user-42",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
cache_key="jwt_key_mapping:sub:user-42",
)
# Caller still receives the winner's mapping even when cleanup fails
assert result == mock_key_obj
prisma_client.db.litellm_verificationtoken.delete.assert_called_once()
@pytest.mark.asyncio
async def test_auto_register_raises_503_when_winner_mapping_vanishes():
"""
Race edge case: this request loses the unique-constraint race, deletes its
orphan, then refetches the winner's mapping — but the winner's row was
concurrently deleted. Previously this returned None, silently falling
through to less-restrictive team-based JWT auth (bypassing the configured
AUTO_REGISTER policy). Must now raise HTTP 503 so the caller retries
rather than getting unintended fallback access.
"""
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(
side_effect=Exception("Unique constraint failed (P2002)")
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock()
# Winner row no longer exists by the time we refetch
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
user_api_key_cache = DualCache()
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": "sk-loser", "key": "sk-loser"},
),
pytest.raises(HTTPException) as exc_info,
):
await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="user-42",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
cache_key="jwt_key_mapping:sub:user-42",
)
assert exc_info.value.status_code == 503
assert "concurrently removed" in exc_info.value.detail
@pytest.mark.asyncio
async def test_proxy_admin_sentinel_skips_db_lookup_on_cache_hit():
"""
When the cache holds the proxy-admin sentinel (written after a prior
request's is_proxy_admin early-return), _resolve_jwt_to_virtual_key must
return None *without* hitting the DB. Caller proceeds to auth_builder.
Without this, every subsequent proxy-admin request under AUTO_REGISTER
would re-query get_jwt_key_mapping_object — a cache-miss regression
introduced by the deferred-auto-register refactor.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
virtual_key_mapping_cache_ttl=300,
)
jwt_claims = {"sub": "admin-user"}
prisma_client = MagicMock()
# Will fail the test if accessed — proves the sentinel short-circuits DB
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(
side_effect=AssertionError("DB must not be hit when sentinel is cached")
)
user_api_key_cache = DualCache()
await user_api_key_cache.async_set_cache(
"jwt_key_mapping:sub:admin-user", "__JWT_PROXY_ADMIN__"
)
result = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert result is None
prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called()
# ──────────────────────────────────────────────
# Tests: AUTO_REGISTER stamps validated identity from auth_builder
# ──────────────────────────────────────────────
@pytest.mark.asyncio
async def test_auto_register_helper_stamps_validated_identity_context():
"""
The deferred-auto-register contract: _auto_register_jwt_mapping is called
with identity fields from JWTAuthManager.auth_builder's *validated*
result (after RBAC, scope mappings, custom_validate, email-domain policy).
These must be passed to generate_key_helper_fn so the created key carries
them — the cached future-request path then inherits the same team/user/org
limits the auth_builder path would have applied.
"""
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
virtual_key_mapping_cache_ttl=300,
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
mock_key_obj = UserAPIKeyAuth(
token="hashed", team_id="validated-team", user_id="validated-user"
)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
) as mock_get_key,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_gen_key,
):
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
mock_get_key.return_value = mock_key_obj
result = await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="new-user",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=DualCache(),
parent_otel_span=None,
proxy_logging_obj=None,
cache_key="jwt_key_mapping:sub:new-user",
team_id="validated-team",
user_id="validated-user",
org_id="validated-org",
end_user_id="validated-end-user",
)
assert result == mock_key_obj
assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team"
assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user"
assert mock_gen_key.call_args.kwargs["organization_id"] == "validated-org"
assert result.org_id == "validated-org"
assert result.end_user_id == "validated-end-user"
# ──────────────────────────────────────────────
# Tests: backward-compat alias jwt_client_id_field
# ──────────────────────────────────────────────
def test_jwt_client_id_field_alias_maps_to_virtual_key_claim_field():
"""
jwt_client_id_field (old doc name) must silently alias to virtual_key_claim_field.
"""
auth = LiteLLM_JWTAuth(jwt_client_id_field="azp")
assert auth.virtual_key_claim_field == "azp"
def test_jwt_client_id_field_does_not_raise_on_duplicate():
"""
If both jwt_client_id_field and virtual_key_claim_field are supplied,
virtual_key_claim_field takes precedence and no error is raised.
"""
auth = LiteLLM_JWTAuth(
jwt_client_id_field="old_field",
virtual_key_claim_field="new_field",
)
assert auth.virtual_key_claim_field == "new_field"

View file

@ -32,7 +32,7 @@ logging.basicConfig(
format="%(asctime)s - %(levelname)s - %(message)s",
)
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import FastAPI
@ -1123,6 +1123,14 @@ from litellm.proxy.management_endpoints.team_endpoints import team_member_add
from test_key_generate_prisma import prisma_client
@pytest.fixture
def mock_prisma_client():
client = MagicMock()
client.connect = AsyncMock()
client.disconnect = AsyncMock()
return client
@pytest.mark.skip(reason="Requires reliable external DB connection (prisma).")
@pytest.mark.parametrize(
"user_role",
@ -1289,7 +1297,6 @@ async def test_create_team_member_add_team_admin_user_api_key_auth(
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
setattr(litellm, "max_internal_user_budget", 10)
setattr(litellm, "internal_user_budget_duration", "5m")
await litellm.proxy.proxy_server.prisma_client.connect()
user = f"ishaan {uuid.uuid4().hex}"
_team_id = "litellm-test-client-id-new"
user_key = "sk-12345678"
@ -1364,7 +1371,6 @@ async def test_create_team_member_add_team_admin(
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
setattr(litellm, "max_internal_user_budget", 10)
setattr(litellm, "internal_user_budget_duration", "5m")
await litellm.proxy.proxy_server.prisma_client.connect()
user = f"ishaan {uuid.uuid4().hex}"
_team_id = "litellm-test-client-id-new"
user_key = "sk-12345678"
@ -1605,7 +1611,10 @@ async def test_add_callback_via_key(prisma_client):
],
)
async def test_add_callback_via_key_litellm_pre_call_utils(
prisma_client, callback_type, expected_success_callbacks, expected_failure_callbacks
mock_prisma_client,
callback_type,
expected_success_callbacks,
expected_failure_callbacks,
):
import json
@ -1614,9 +1623,8 @@ async def test_add_callback_via_key_litellm_pre_call_utils(
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
proxy_config = getattr(litellm.proxy.proxy_server, "proxy_config")
@ -1762,7 +1770,10 @@ async def test_disable_fallbacks_by_key(disable_fallbacks_set):
],
)
async def test_add_callback_via_key_litellm_pre_call_utils_gcs_bucket(
prisma_client, callback_type, expected_success_callbacks, expected_failure_callbacks
mock_prisma_client,
callback_type,
expected_success_callbacks,
expected_failure_callbacks,
):
import json
@ -1771,9 +1782,8 @@ async def test_add_callback_via_key_litellm_pre_call_utils_gcs_bucket(
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
proxy_config = getattr(litellm.proxy.proxy_server, "proxy_config")
@ -1896,7 +1906,10 @@ async def test_add_callback_via_key_litellm_pre_call_utils_gcs_bucket(
],
)
async def test_add_callback_via_key_litellm_pre_call_utils_langsmith(
prisma_client, callback_type, expected_success_callbacks, expected_failure_callbacks
mock_prisma_client,
callback_type,
expected_success_callbacks,
expected_failure_callbacks,
):
import json
@ -1905,9 +1918,8 @@ async def test_add_callback_via_key_litellm_pre_call_utils_langsmith(
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
proxy_config = getattr(litellm.proxy.proxy_server, "proxy_config")

View file

@ -31,12 +31,14 @@ from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_PendingAutoRegister,
_matches_routing_override,
_reserve_budget_after_common_checks,
_route_requires_auth_despite_public,
_routing_selector_matches_claim,
_run_centralized_common_checks,
_run_post_custom_auth_checks,
_user_api_key_auth_builder,
get_api_key,
user_api_key_auth,
)
@ -1550,6 +1552,93 @@ class TestJWTOAuth2Coexistence:
assert mock_jwt_auth.call_args.kwargs["request_method"] == "POST"
assert result.user_id == "jwt-human-user"
@pytest.mark.asyncio
async def test_auto_register_passes_validated_org_context_to_generated_key(self):
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"})
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
virtual_key_mapping_cache_ttl=300,
)
auto_registered_key = UserAPIKeyAuth(
token="hashed-auto-key",
team_id="validated-team",
user_id="validated-user",
org_id="validated-org",
end_user_id="validated-end-user",
)
mock_jwt_result = {
"is_proxy_admin": False,
"team_object": None,
"user_object": None,
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": "validated-team",
"user_id": "validated-user",
"end_user_id": "validated-end-user",
"org_id": "validated-org",
"team_membership": None,
"jwt_claims": {"sub": "user1"},
}
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.method = "POST"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
mock_request.state = SimpleNamespace()
with (
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler),
patch(
"litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key",
new_callable=AsyncMock,
return_value=_PendingAutoRegister(
claim_field="sub",
claim_value="user1",
cache_key="jwt_key_mapping:sub:user1",
),
),
patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=mock_jwt_result,
),
patch(
"litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping",
new_callable=AsyncMock,
return_value=auto_registered_key,
) as mock_auto_register,
):
result = await _user_api_key_auth_builder(
request=mock_request,
api_key=jwt_token,
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o-mini"},
)
mock_auto_register.assert_awaited_once()
assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team"
assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user"
assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org"
assert mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user"
assert result.org_id == "validated-org"
@pytest.mark.asyncio
async def test_routing_override_routes_matching_jwt_to_oauth2(self):
"""

View file

@ -4,86 +4,11 @@
"count": 1
}
},
"src/app/(dashboard)/hooks/accessGroups/useAccessGroupDetails.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/accessGroups/useAccessGroups.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/accessGroups/useCreateAccessGroup.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/accessGroups/useDeleteAccessGroup.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/accessGroups/useEditAccessGroup.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/blogPosts/useBlogPosts.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/cloudzero/useCloudZeroCreate.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/cloudzero/useCloudZeroDryRun.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/cloudzero/useCloudZeroExport.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/cloudzero/useCloudZeroSettings.ts": {
"no-restricted-syntax": {
"count": 3
}
},
"src/app/(dashboard)/hooks/configOverrides/hashicorpVaultApi.ts": {
"no-restricted-syntax": {
"count": 4
}
},
"src/app/(dashboard)/hooks/guardrails/useRegisterGuardrail.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/healthReadiness/useHealthReadinessDetails.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/keys/useKeyAliases.test.ts": {
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/hooks/keys/useKeys.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/keys/useResetKeySpend.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/models/useModels.ts": {
"max-params": {
"count": 1
@ -94,76 +19,26 @@
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useCreateProject.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useDeleteProject.test.ts": {
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useDeleteProject.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useProjectDetails.test.ts": {
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useProjectDetails.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useProjects.test.ts": {
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useProjects.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useUpdateProject.test.ts": {
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/hooks/projects/useUpdateProject.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/proxyConfig/useProxyConfig.ts": {
"no-restricted-syntax": {
"count": 2
}
},
"src/app/(dashboard)/hooks/router/useRouterFields.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/storeModelInDB/useStoreModelInDB.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/storeRequestInSpendLogs/useStoreRequestInSpendLogs.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/app/(dashboard)/hooks/teams/useTeams.ts": {
"no-restricted-syntax": {
"count": 2
}
},
"src/app/(dashboard)/layout.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -356,11 +231,6 @@
"count": 1
}
},
"src/components/CostTrackingSettings/pricing_calculator/use_multi_cost_estimate.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/CostTrackingSettings/provider_discount_table.test.tsx": {
"unused-imports/no-unused-imports": {
"count": 1
@ -381,16 +251,6 @@
"count": 1
}
},
"src/components/CostTrackingSettings/use_discount_config.ts": {
"no-restricted-syntax": {
"count": 2
}
},
"src/components/CostTrackingSettings/use_margin_config.ts": {
"no-restricted-syntax": {
"count": 2
}
},
"src/components/CreateUserButton.tsx": {
"no-restricted-imports": {
"count": 1
@ -702,9 +562,6 @@
}
},
"src/components/WebRTCTester.jsx": {
"no-restricted-syntax": {
"count": 2
},
"react/no-unescaped-entities": {
"count": 2
}
@ -973,9 +830,6 @@
"no-restricted-imports": {
"count": 1
},
"no-restricted-syntax": {
"count": 3
},
"react-hooks/immutability": {
"count": 1
}
@ -1178,9 +1032,6 @@
"no-restricted-imports": {
"count": 1
},
"no-restricted-syntax": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1253,11 +1104,6 @@
"count": 1
}
},
"src/components/mcp_tools/ByokCredentialModal.tsx": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/mcp_tools/MCPLogoSelector.test.tsx": {
"unused-imports/no-unused-imports": {
"count": 1
@ -1502,9 +1348,6 @@
"src/components/networking.tsx": {
"max-params": {
"count": 23
},
"no-restricted-syntax": {
"count": 273
}
},
"src/components/object_permissions_view.tsx": {
@ -1621,11 +1464,6 @@
"count": 13
}
},
"src/components/playground/chat_ui/CodeInterpreterOutput.tsx": {
"no-restricted-syntax": {
"count": 2
}
},
"src/components/playground/chat_ui/CodeInterpreterTool.tsx": {
"no-restricted-imports": {
"count": 1
@ -1657,9 +1495,6 @@
"src/components/playground/llm_calls/a2a_send_message.tsx": {
"max-params": {
"count": 2
},
"no-restricted-syntax": {
"count": 2
}
},
"src/components/playground/llm_calls/anthropic_messages.tsx": {
@ -1685,14 +1520,6 @@
"src/components/playground/llm_calls/embeddings_api.tsx": {
"max-params": {
"count": 1
},
"no-restricted-syntax": {
"count": 1
}
},
"src/components/playground/llm_calls/fetch_agents.tsx": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/playground/llm_calls/image_edits.tsx": {
@ -1708,9 +1535,6 @@
"src/components/playground/llm_calls/interactions_api.tsx": {
"max-params": {
"count": 1
},
"no-restricted-syntax": {
"count": 1
}
},
"src/components/playground/llm_calls/responses_api.tsx": {
@ -1904,11 +1728,6 @@
"count": 1
}
},
"src/components/prompts/prompt_editor_view/conversation_panel/useConversation.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/prompts/prompt_info.tsx": {
"no-restricted-imports": {
"count": 1
@ -1970,11 +1789,6 @@
"count": 1
}
},
"src/components/survey/SurveyModal.tsx": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/tag_management/TagTable.tsx": {
"no-restricted-imports": {
"count": 1
@ -2037,11 +1851,6 @@
"count": 1
}
},
"src/components/team/useMyTeamMember.ts": {
"no-restricted-syntax": {
"count": 1
}
},
"src/components/templates/key_edit_view.tsx": {
"no-restricted-imports": {
"count": 1
@ -2069,9 +1878,6 @@
"no-restricted-imports": {
"count": 1
},
"no-restricted-syntax": {
"count": 3
},
"react-hooks/immutability": {
"count": 1
}
@ -2223,9 +2029,6 @@
}
},
"src/components/workflow_runs/index.tsx": {
"no-restricted-syntax": {
"count": 3
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -2235,11 +2038,6 @@
"count": 1
}
},
"src/contexts/ThemeContext.tsx": {
"no-restricted-syntax": {
"count": 1
}
},
"src/data/claimsCompliancePrompts.ts": {
"max-params": {
"count": 1

View file

@ -32,13 +32,6 @@ const eslintConfig = [
"max-depth": ["warn", 4],
"max-params": ["error", 4],
"max-nested-callbacks": ["error", 4],
"no-restricted-syntax": [
"error",
{
selector: "CallExpression[callee.name='fetch']",
message: "Use React Query (@tanstack/react-query) for data fetching instead of a raw fetch().",
},
],
"no-restricted-imports": [
"error",
{