feat(jwt): auto_register_map_existing_key maps JWT to the user's existing virtual key (#42375)

* test(e2e): jwt auto_register map-existing-key repro

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(jwt): auto_register_map_existing_key maps JWT to the user's existing virtual key

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(jwt): exclude blocked keys from auto_register_map_existing_key reuse

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(jwt): route existing-key lookup through VerificationTokenRepository

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(e2e): stop requiring LITELLM_SALT_KEY for the owned JWT gateway

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(e2e): gate the owned JWT gateway tests behind E2E_OWNED_GATEWAY

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(jwt): only reuse keys that can call LLM routes in auto_register_map_existing_key

Skip Admin UI session keys and keys whose allowed_routes restrict them to
anything other than llm_api_routes (management, read_only, password-reset
sessions). Mapping a JWT to one of those left the user with 401s or 403s on
every LLM call, since the mapping persists.

* fix(jwt): scope auto_register_map_existing_key reuse to the JWT-resolved team

Only reuse a key whose team_id matches the team auth_builder resolved for
the JWT (no team matches no team), so a personal key can no longer bypass
the resolved team's model and budget limits.

With the flag on, the first JWT request now falls through to the same
virtual-key checks later mapped requests get, instead of returning early,
so a reused key's own limits apply from request one rather than 200 then
403. Flag off keeps the early return unchanged.

* fix(jwt): keep the early return when no master key is set

Without a master key the generic virtual-key path returns a bare
INTERNAL_USER object, so falling through on the first auto-registered
request dropped the key's team, models and budgets. Only fall through when
a master key is configured.

Tests now assert the reused key per team rather than the query shape, and
cover the flag-off early return and the no-master-key case.

* test(jwt): assert on race-loser's returned key, not only mocks (TQ002)

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* fix(jwt): close the auto_register_map_existing_key race, shared-claim and expiry holes

A key auto_register just minted is never adopted by a concurrent request, so the race loser's cleanup can no longer delete a key another request mapped and cascade its mapping away (503, user left with no key)

Reuse only happens when the claim value is the JWT-resolved user_id. A shared claim such as azp or client_id falls back to minting, so one user can no longer land on another user's personal key and budget

Only keys that never expire are reused, so an expiring key can no longer pin the claim to a permanent 401

Integration tests on a real proxy and Postgres cover all three. The race test holds the first mapping insert in a Postgres relay, so the interleaving is forced rather than timed. The where-clause shape unit tests are replaced by these, since only a real database proves the filter

* test(e2e): create the reused key in the team the JWT resolves to

The flag only reuses a key in the JWT-resolved team, and this identity's groups claim resolves to its team, so a teamless key was never eligible and the test could not pass

* test(integration): match the held statement across TCP reads

The relay looked for the trigger inside one read, so an insert split across two reads was never held and the race test would fail waiting for it. It now matches one exact trigger over a window that keeps the end of the previous read

* fix(jwt): gate key reuse on the claim field, not on the claim value

Requiring the claim value to equal the resolved user_id skipped reuse for users matched through the sso_user_id or case-insensitive email fallback, whose stored user_id differs from the JWT sub. That is the lookup LIT-5378 asks for. Reuse is now allowed when the virtual key claim is the user_id or user_email JWT field, globally or for the token's issuer, which still keeps shared claims such as azp or client_id on the mint path

* fix(jwt): let an issuer's own user field replace the global one when gating key reuse

An issuer that identifies users by uid no longer treats the global sub field as a user identity claim, so a shared sub under that issuer mints instead of reusing a personal key

* test(jwt): make the flag-off test fail when the flag no longer gates key reuse

The flag-off test used a config where sub was not a user identity claim, so deleting the flag check still passed. Configure user_id_jwt_field=sub so only the flag keeps the lookup off, and drop test docstrings

* chore(lint): drop mutable-ok suppressions that LIT013 flags as no-ops

---------

Co-authored-by: yuneng <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mrinal <mrinal@berri.ai>
Co-authored-by: Mrinal Chanshetty <mchanshetty@Mrinals-MacBook-Pro.local>
Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-02 12:11:33 -07:00 • committed by GitHub
parent 2584721ca3
commit b21e44cbf9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 1333 additions and 111 deletions

View file

@ -176,6 +176,7 @@ jobs:
TESTS: ${{ needs.detect.outputs.tests }}
E2E_FIXTURE_MODE: live
E2E_PROVIDER_EDGE_HOST_REACHABLE: '1'
E2E_OWNED_GATEWAY: '1'
COLUMNS: '400'
run: |
umask 077

View file

@ -5464,6 +5464,17 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"'auto_register': auto-create a virtual key and mapping on first encounter."
),
)
auto_register_map_existing_key: bool = Field(
default=False,
description=(
"Only used with unregistered_jwt_client_behavior='auto_register'. When True and the virtual key claim "
"field is the user_id_jwt_field or user_email_jwt_field, the JWT claim is mapped to a virtual key the "
"JWT-resolved user already owns instead of minting a new one. If the user owns several, the most recently created key in the "
"JWT-resolved team (or with no team when the JWT resolves none) is chosen among keys that never "
"expire, are not blocked, are not Admin UI session keys, were not minted by auto_register, and "
"have no allowed_routes or include llm_api_routes. Otherwise a new key is minted as usual."
),
)
routing_overrides: list[JWTRoutingOverride] | None = Field(
default=None,
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
@ -5564,6 +5575,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
return issuer_config.virtual_key_claim_field
return self.virtual_key_claim_field
def is_user_identity_claim(self, claim_field: str, issuer: str | None) -> bool:
issuer_config: Final = self.get_issuer_config(issuer)
if issuer_config is None:
return claim_field in (self.user_id_jwt_field, self.user_email_jwt_field)
return claim_field in (
issuer_config.user_id_jwt_field or self.user_id_jwt_field,
issuer_config.user_email_jwt_field or self.user_email_jwt_field,
)
def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior:
issuer_config: Final = self.get_issuer_config(issuer)
if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None:

View file

@ -140,6 +140,7 @@ from litellm.proxy.utils import (
normalize_route_for_root_path,
)
from litellm.repositories.table_repositories import TeamMembershipRepository
from litellm.repositories.verification_token_repository import VerificationTokenRepository
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.secret_managers.main import get_secret_bool
from litellm.types.services import ServiceTypes
@ -939,6 +940,24 @@ class _PendingAutoRegister(NamedTuple):
jwt_issuer: str | None = None
def _claim_identifies_user(jwt_handler: JWTHandler, claim_field: str, jwt_issuer: str | None) -> bool:
if not jwt_handler.litellm_jwtauth.auto_register_map_existing_key:
return False
if jwt_handler.litellm_jwtauth.is_user_identity_claim(claim_field, jwt_issuer):
return True
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register_map_existing_key): claim '%s' is not the user_id or user_email JWT field "
"and may be shared by several users, so a new key is minted instead of reusing one the user owns.",
claim_field,
)
return False
async def _reusable_key_hash_for_user(prisma_client: PrismaClient, user_id: str, team_id: str | None) -> str | None:
key: Final = await VerificationTokenRepository(prisma_client).find_newest_reusable_llm_api_key(user_id, team_id)
return None if key is None else key.token
async def _auto_register_jwt_mapping(
virtual_key_claim_field: str,
claim_value: str,
@ -957,8 +976,10 @@ async def _auto_register_jwt_mapping(
) -> UserAPIKeyAuth | None:
"""
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
claim value, or point the mapping at a key the resolved user already owns
when ``auto_register_map_existing_key`` is set. ``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.
@ -974,29 +995,38 @@ async def _auto_register_jwt_mapping(
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: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
table_name="key",
team_id=team_id,
user_id=user_id,
organization_id=org_id,
agent_id=agent_id,
metadata={
"auto_registered": True,
"jwt_claim_field": virtual_key_claim_field,
"jwt_claim_value": claim_value,
},
existing_token_hash: Final = (
await _reusable_key_hash_for_user(prisma_client, user_id, team_id)
if user_id is not None and _claim_identifies_user(jwt_handler, virtual_key_claim_field, jwt_issuer)
else None
)
# 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"])
minted: Final = existing_token_hash is None
if existing_token_hash is not None:
token_hash = existing_token_hash
else:
# ``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: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
table_name="key",
team_id=team_id,
user_id=user_id,
organization_id=org_id,
agent_id=agent_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(
@ -1023,15 +1053,16 @@ async def _auto_register_jwt_mapping(
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,
)
if minted:
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,
@ -1061,7 +1092,8 @@ async def _auto_register_jwt_mapping(
)
verbose_proxy_logger.info(
"JWT Key Mapping (auto_register): created new virtual key for %s='%s'.",
"JWT Key Mapping (auto_register): %s virtual key for %s='%s'.",
"created new" if minted else "mapped existing",
virtual_key_claim_field,
claim_value,
)
@ -1075,7 +1107,8 @@ async def _auto_register_jwt_mapping(
).resolve(hashed_token=token_hash)
)
if auto_registered_key is not None:
auto_registered_key.org_id = org_id
if minted:
auto_registered_key.org_id = org_id
auto_registered_key.end_user_id = end_user_id
auto_registered_key.api_key = auto_registered_key.token
return auto_registered_key
@ -1771,8 +1804,8 @@ async def _user_api_key_auth_builder(
# 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: Final = await _auto_register_jwt_mapping(
auto_registered: Final = (
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,
@ -1788,72 +1821,81 @@ async def _user_api_key_auth_builder(
end_user_id=end_user_id,
agent_id=agent_id,
)
if auto_registered is not None:
auto_registered.jwt_claims = jwt_claims
auto_registered.user_email = user_email
# The auto-registered token is built from the new key's
# columns, which carry no user budget. Carry over the
# already-loaded user row rather than re-reading it, or
# the budget check below has nothing to enforce.
auto_registered.user_model_max_budget = (
user_object.model_max_budget if user_object is not None else None
)
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,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
if pending_auto_register is not None and prisma_client is not None
else None
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
# Fetch project object for JWT path if project_id is set
_jwt_project_obj = None
if valid_token.project_id is not None:
_jwt_project_obj = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
if auto_registered is not None:
auto_registered.jwt_claims = jwt_claims
auto_registered.user_email = user_email
# The auto-registered token is built from the new key's
# columns, which carry no user budget. Carry over the
# already-loaded user row rather than re-reading it, or
# the budget check below has nothing to enforce.
auto_registered.user_model_max_budget = (
user_object.model_max_budget if user_object is not None else None
)
if _jwt_project_obj is not None:
valid_token.project_metadata = _jwt_project_obj.metadata
valid_token.project_alias = _jwt_project_obj.project_alias
valid_token = auto_registered
api_key = valid_token.token or ""
# JWT auth returns here rather than falling through to the
# virtual-key checks below, so the user's per-model budget
# has to be enforced on this path too. Without it the
# post-call increment still charges the counter and nothing
# ever reads it, which is worse than not tracking at all.
# Guarded by the same flag the virtual-key path uses, or a
# zero-cost model would be refused here and allowed there,
# while the log above claims all budget checks were skipped.
if not skip_budget_checks:
await _check_user_model_budget(
valid_token=cast(UserAPIKeyAuth, valid_token),
model_max_budget_limiter=model_max_budget_limiter,
models=_get_model_names_for_budget_checks(
model=_get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
)
),
falls_through_to_key_checks: Final = (
auto_registered is not None
and jwt_handler.litellm_jwtauth.auto_register_map_existing_key
and master_key is not None
)
if not falls_through_to_key_checks:
# Check if model has zero cost - if so, skip all budget checks
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
return cast(UserAPIKeyAuth, valid_token)
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
# Fetch project object for JWT path if project_id is set
_jwt_project_obj = None
if valid_token.project_id is not None:
_jwt_project_obj = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if _jwt_project_obj is not None:
valid_token.project_metadata = _jwt_project_obj.metadata
valid_token.project_alias = _jwt_project_obj.project_alias
# JWT auth returns here rather than falling through to the
# virtual-key checks below, so the user's per-model budget
# has to be enforced on this path too. Without it the
# post-call increment still charges the counter and nothing
# ever reads it, which is worse than not tracking at all.
# Guarded by the same flag the virtual-key path uses, or a
# zero-cost model would be refused here and allowed there,
# while the log above claims all budget checks were skipped.
if not skip_budget_checks:
await _check_user_model_budget(
valid_token=cast(UserAPIKeyAuth, valid_token),
model_max_budget_limiter=model_max_budget_limiter,
models=_get_model_names_for_budget_checks(
model=_get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
)
),
)
return cast(UserAPIKeyAuth, valid_token)
#### ELSE ####
## CHECK PASS-THROUGH ENDPOINTS ##

View file

@ -8,6 +8,7 @@ from datetime import datetime
from types import TracebackType
from typing import TYPE_CHECKING, Final, Protocol
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.models.verification_token import (
LiteLLM_VerificationToken,
)
@ -123,6 +124,37 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id})
return self._to_model_list(records)
async def find_newest_reusable_llm_api_key(
self, user_id: str, team_id: str | None
) -> LiteLLM_VerificationToken | None:
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(
where={
"user_id": user_id,
"team_id": team_id,
"expires": None,
"AND": [
{"OR": [{"blocked": False}, {"blocked": None}]},
{
"OR": [
{"team_id": None},
{"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}},
]
},
{
"OR": [
{"allowed_routes": {"is_empty": True}},
{"allowed_routes": {"has": "llm_api_routes"}},
]
},
],
},
order={"created_at": "desc"},
)
return next(
(key for key in self._to_model_list(records) if key.metadata.get("auto_registered") is not True),
None,
)
async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]:
"""Find all tokens belonging to a team."""
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id})

View file

@ -32,6 +32,7 @@ from e2e_config import (
MCP_OAUTH_LIVE_OPT_IN_ENV,
OTEL_TLS_OPT_IN_ENV,
OTEL_V2_OPT_IN_ENV,
OWNED_GATEWAY_OPT_IN_ENV,
PROMPT_CACHING_OPT_IN_ENV,
PROVIDER_EDGE_HOST_OPT_IN_ENV,
PROXY_BASE_URL,
@ -70,6 +71,7 @@ OPT_IN_MARKERS: Final = MappingProxyType(
"cli_determinism": CLI_DETERMINISM_OPT_IN_ENV,
"mcp_oauth_live": MCP_OAUTH_LIVE_OPT_IN_ENV,
"provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV,
"owned_gateway": OWNED_GATEWAY_OPT_IN_ENV,
"otel_v2": OTEL_V2_OPT_IN_ENV,
"otel_tls": OTEL_TLS_OPT_IN_ENV,
"secret_manager": SECRET_MANAGER_OPT_IN_ENV,
@ -172,6 +174,11 @@ def pytest_configure(config: pytest.Config) -> None:
"provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the "
"gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set",
)
config.addinivalue_line(
"markers",
"owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL "
"on the pytest host; deselected unless E2E_OWNED_GATEWAY is set",
)
config.addinivalue_line(
"markers",
"otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set",

View file

@ -60,6 +60,9 @@
- {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"}
- {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"}
- {id: other.auth.jwt.auto_register_maps_existing_key, module: other, tier: P0, area: auth, assertions: [maps_existing_key], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, the first JWT call of a user who already owns a key writes the sub-claim mapping to that existing key hash and mints nothing; the spend row lands on the pre-existing key (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.jwt.auto_register_mints_when_keyless, module: other, tier: P0, area: auth, assertions: [mints_when_keyless], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, a user with no keys still gets exactly one minted key and a sub-claim mapping on their first JWT call (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.jwt.auto_register_default_mints, module: other, tier: P0, area: auth, assertions: [default_mints], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "Without auto_register_map_existing_key, auto_register keeps the current behavior: it mints a second key for a user who already has one and bills the minted key (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"}
- {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"}
- {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"}

View file

@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy.
from __future__ import annotations
import os
import socket
from dataclasses import dataclass
import time
import uuid
@ -16,6 +17,7 @@ from typing import Final
from dotenv import load_dotenv
from fixture_mode import deterministic_marker, parse_fixture_mode, registration_owner
from provider_edge import provider_edge_api_base
from pydantic import TypeAdapter
# Local runs keep provider / DataDog keys in tests/e2e/.env (see CONTRIBUTING.md).
# Compose injects them into the proxy container, but pytest on the host does not
@ -206,6 +208,7 @@ REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS"
CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM"
MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE"
PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE"
OWNED_GATEWAY_OPT_IN_ENV: Final = "E2E_OWNED_GATEWAY"
OTEL_V2_OPT_IN_ENV: Final = "E2E_OTEL_V2"
OTEL_TLS_OPT_IN_ENV: Final = "E2E_OTEL_EXPORTER_ENDPOINT"
SECRET_MANAGER_OPT_IN_ENV: Final = "E2E_SECRET_MANAGER"
@ -296,6 +299,15 @@ def unique_marker() -> str:
return uuid.uuid4().hex[:12]
INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_")
def available_port() -> int:
with socket.socket() as listener:
listener.bind(("127.0.0.1", 0))
return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1]
def settle_propagation(written_at: float) -> None:
"""Block until PROPAGATION_TIMEOUT has elapsed since `written_at`, a
`time.monotonic()` stamp taken the moment a control-plane write returned.

View file

@ -8,7 +8,6 @@ The optional live edge measures headers without recording credentials or bodies.
from __future__ import annotations
import os
import socket
import subprocess
import sys
import threading
@ -20,13 +19,12 @@ from pathlib import Path
from typing import Final
import psycopg
from e2e_config import INHERITED_ENV_PREFIXES, available_port
from e2e_http import NoBody
from idp import Keycloak, stop_process_group
from proxy_client import ProxyClient, build_proxy_client
from psycopg.rows import class_row
from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError
INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_")
from pydantic import BaseModel, SecretStr, ValidationError
class StoredOAuth(BaseModel):
@ -101,12 +99,6 @@ class OAuthObservation:
assert all(not item[2] for item in snapshot), "gateway bearer leaked to the upstream"
def available_port() -> int:
with socket.socket() as listener:
listener.bind(("127.0.0.1", 0))
return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1]
@dataclass(slots=True)
class OAuthGateway:
base_url: str

View file

@ -1537,6 +1537,7 @@ class UserNewBody(BaseModel):
class UserNewResponse(BaseModel):
user_id: str
key: str | None = None
class UserUpdateBody(BaseModel):
@ -1580,6 +1581,40 @@ class UserListResponse(BaseModel):
total: int
class UserKeyRow(BaseModel):
token: str
key_alias: str | None = None
class UserInfoWithKeysResponse(BaseModel):
user_id: str | None = None
keys: list[UserKeyRow] = []
class JwtKeyMappingRow(BaseModel):
id: str
jwt_claim_name: str
jwt_claim_value: str
created_by: str | None = None
class JwtKeyMappingListParams(BaseModel):
size: int = 100
class JwtKeyMappingListResponse(BaseModel):
mappings: list[JwtKeyMappingRow]
total_count: int
class JwtKeyMappingDeleteBody(BaseModel):
id: str
class JwtKeyMappingDeleteResponse(BaseModel):
status: str
class OrgNewBody(BaseModel):
organization_alias: str
models: list[str] = []

View file

@ -21,12 +21,20 @@ from idp import Keycloak, keycloak_from_env
from models import (
ChatBody,
ChatResponse,
JwtKeyMappingDeleteBody,
JwtKeyMappingDeleteResponse,
JwtKeyMappingListParams,
JwtKeyMappingListResponse,
ModelsListParams,
ModelsListResponse,
ReadinessDetailsResponse,
ReadinessResponse,
UserInfoParams,
UserInfoWithKeysResponse,
UserListParams,
UserListResponse,
UserNewBody,
UserNewResponse,
)
from proxy_client import ProxyClient
from pydantic import Field
@ -79,6 +87,44 @@ class OtherClient:
response_type=ReadinessDetailsResponse,
)
def user_new(self, body: UserNewBody) -> Result[UserNewResponse]:
"""POST /user/new under the master key: seed the litellm user a JWT
`sub` claim resolves to, before that token ever reaches the proxy."""
return self.proxy.transport.post(
"/user/new",
headers=self.proxy.transport.master,
json=body,
response_type=UserNewResponse,
)
def user_info(self, user_id: str) -> Result[UserInfoWithKeysResponse]:
"""GET /user/info under the master key. Only the user's key rows are
modelled: `token` is the stored key hash, never the plaintext key."""
return self.proxy.transport.get(
"/user/info",
headers=self.proxy.transport.master,
params=UserInfoParams(user_id=user_id),
response_type=UserInfoWithKeysResponse,
)
def jwt_mapping_list(self) -> Result[JwtKeyMappingListResponse]:
"""GET /jwt/key/mapping/list under the master key."""
return self.proxy.transport.get(
"/jwt/key/mapping/list",
headers=self.proxy.transport.master,
params=JwtKeyMappingListParams(size=100),
response_type=JwtKeyMappingListResponse,
)
def jwt_mapping_delete(self, mapping_id: str) -> Result[JwtKeyMappingDeleteResponse]:
"""POST /jwt/key/mapping/delete under the master key."""
return self.proxy.transport.post(
"/jwt/key/mapping/delete",
headers=self.proxy.transport.master,
json=JwtKeyMappingDeleteBody(id=mapping_id),
response_type=JwtKeyMappingDeleteResponse,
)
def chat_as_team(self, token: str, team: str, body: ChatBody) -> Result[ChatResponse]:
"""POST /chat/completions under `token` with `x-litellm-team-id: team`."""
return self.proxy.transport.post(

View file

@ -0,0 +1,108 @@
"""An owned, source-built proxy whose `litellm_jwtauth` block a test controls.
The shared proxy on :4000 runs the CONTRIBUTING.md JWT block, so a test that
needs a different `litellm_jwtauth` config boots its own gateway on a free port
against the same database and the same Keycloak realm. The caller supplies the
`litellm_jwtauth` mapping verbatim, which is exactly what makes a config an
unfixed proxy rejects observable as a boot failure in this gateway's own log.
"""
from __future__ import annotations
import os
import subprocess
import sys
import time
from collections.abc import Mapping
from contextlib import ExitStack
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
from e2e_config import INHERITED_ENV_PREFIXES, available_port
from e2e_http import NoBody
from idp import Keycloak, stop_process_group
from proxy_client import ProxyClient, build_proxy_client
MODEL_NAME: Final = "gemini-3.8-flash"
@dataclass(slots=True)
class OwnedJwtGateway:
base_url: str
proxy: ProxyClient
_environment: Mapping[str, str] = field(repr=False)
_command: tuple[str, ...] = field(repr=False)
_log_path: Path
_child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False)
def start(self) -> None:
with self._log_path.open("ab") as log:
self._child = subprocess.Popen(
self._command,
env=self._environment,
stdout=log,
stderr=log,
start_new_session=True,
)
deadline: Final = time.monotonic() + 120
while time.monotonic() < deadline:
assert self._child.poll() is None, "owned JWT gateway exited; inspect its private log"
result = self.proxy.transport.probe("/health/liveliness", params=NoBody())
if result.status_code == 200:
return
time.sleep(0.5)
raise AssertionError("owned JWT gateway did not become ready")
def stop(self) -> None:
if self._child is not None:
stop_process_group(self._child)
assert self._child.poll() is not None, "old gateway process is still alive"
def owned_jwt_gateway(
idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str
) -> OwnedJwtGateway:
for env_name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_MASTER_KEY"):
assert os.environ.get(env_name), f"{env_name} is required for the owned JWT gateway"
port: Final = available_port()
base_url: Final = f"http://127.0.0.1:{port}"
config: Final = directory / f"{name}.yaml"
config.write_text(
"model_list:\n"
f" - model_name: {MODEL_NAME}\n"
" litellm_params:\n"
f" model: gemini/{MODEL_NAME}\n"
" api_key: os.environ/GEMINI_API_KEY\n"
"general_settings:\n"
" master_key: os.environ/LITELLM_MASTER_KEY\n"
" database_url: os.environ/DATABASE_URL\n"
" proxy_batch_write_at: 5\n"
" enable_jwt_auth: true\n"
" litellm_jwtauth:\n" + "".join(f" {line}\n" for line in litellm_jwtauth.strip().splitlines())
)
environment: Final = {
**{key: value for key, value in os.environ.items() if not key.startswith(INHERITED_ENV_PREFIXES)},
"JWT_PUBLIC_KEY_URL": idp.jwks_url,
"JWT_ISSUER": idp.issuer,
"JWT_AUDIENCE": "litellm-e2e",
"LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true",
"DISABLE_SCHEMA_UPDATE": "true",
"STORE_MODEL_IN_DB": "True",
"PYTHONPATH": str(Path(__file__).resolve().parents[3]),
}
gateway: Final = OwnedJwtGateway(
base_url=base_url,
proxy=build_proxy_client(
base_url=base_url,
control_plane_base_url=base_url,
replica_urls=(base_url,),
master_key=os.environ["LITELLM_MASTER_KEY"],
),
_environment=environment,
_command=(sys.executable, "-m", "litellm.proxy.proxy_cli", "--config", str(config), "--port", str(port)),
_log_path=directory / f"{name}.log",
)
cleanup.callback(gateway.stop)
gateway.start()
return gateway

View file

@ -0,0 +1,182 @@
"""auto_register with auto_register_map_existing_key binds the JWT claim to the user's existing key.
`unregistered_jwt_client_behavior: auto_register` on `virtual_key_claim_field: sub` mints a fresh
virtual key on the user's first JWT call. With `auto_register_map_existing_key: true` the proxy must
instead point the new JWT mapping at a key the resolved user already owns, and mint only when the
user has none. Each behavior gets its own gateway because the flag lives in `litellm_jwtauth`, so
this file boots two owned proxies against the shared database and Keycloak realm.
"""
from __future__ import annotations
import hashlib
from collections.abc import Iterator
from contextlib import ExitStack
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import unwrap
from idp import Identity, Keycloak
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, JwtKeyMappingRow, KeyGenerateBody, TeamNewBody, UserNewBody
from other_client import OtherClient
from owned_jwt_gateway import MODEL_NAME, OwnedJwtGateway, owned_jwt_gateway
pytestmark = pytest.mark.e2e
_JWT_COMMON: Final = (
"user_id_jwt_field: sub\n"
"user_email_jwt_field: email\n"
"team_ids_jwt_field: groups\n"
"user_id_upsert: true\n"
"virtual_key_claim_field: sub\n"
"unregistered_jwt_client_behavior: auto_register"
)
def _key_hash(key: str) -> str:
return hashlib.sha256(key.encode()).hexdigest()
def _ping() -> ChatBody:
return ChatBody(
model=MODEL_NAME,
messages=[ChatMessage(role="user", content=f"Reply with the single word ok. {unique_marker()}")],
max_tokens=5,
)
def _identity_with_user(idp: Keycloak, client: OtherClient, resources: ResourceManager) -> Identity:
"""An IdP identity plus the litellm user and team its claims resolve to, with
teardown that also sweeps the user's keys and JWT mapping rows the proxy
wrote, since those outlive the user row itself."""
marker: Final = unique_marker()
identity: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer)
resources.defer(lambda: client.proxy.delete_user(identity.user_id))
team_id: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-jwt-{marker}", team_id=identity.group))
resources.defer(lambda: client.proxy.delete_team(team_id))
unwrap(
client.user_new(
UserNewBody(
user_id=identity.user_id,
user_email=f"{identity.username}@example.com",
user_role="internal_user",
auto_create_key=False,
)
)
)
def delete_user_keys() -> None:
for row in unwrap(client.user_info(identity.user_id)).keys:
client.proxy.delete_key(row.token)
def delete_user_mappings() -> None:
for mapping in unwrap(client.jwt_mapping_list()).mappings:
if mapping.jwt_claim_value == identity.user_id:
_ = client.jwt_mapping_delete(mapping.id)
resources.defer(delete_user_keys)
resources.defer(delete_user_mappings)
return identity
def _mapping_for(client: OtherClient, claim_value: str) -> JwtKeyMappingRow | None:
return next(
(row for row in unwrap(client.jwt_mapping_list()).mappings if row.jwt_claim_value == claim_value),
None,
)
@pytest.fixture(scope="module")
def mapping_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
with ExitStack() as cleanup:
yield owned_jwt_gateway(
idp,
tmp_path_factory.mktemp("jwt-mapping"),
cleanup,
litellm_jwtauth=f"{_JWT_COMMON}\nauto_register_map_existing_key: true",
name="jwt-mapping-gateway",
)
@pytest.fixture(scope="module")
def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
with ExitStack() as cleanup:
yield owned_jwt_gateway(
idp,
tmp_path_factory.mktemp("jwt-minting"),
cleanup,
litellm_jwtauth=_JWT_COMMON,
name="jwt-minting-gateway",
)
@pytest.mark.owned_gateway
class TestJwtAutoRegisterMapExistingKey:
@pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key")
def test_first_jwt_call_maps_to_the_users_existing_key_and_mints_none(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
existing_key: Final = client.proxy.generate_key(
KeyGenerateBody(
user_id=identity.user_id, team_id=identity.group, key_alias=f"e2e-jwt-existing-{unique_marker()}"
)
)
resources.defer(lambda: client.proxy.delete_key(existing_key))
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert [row.token for row in keys] == [_key_hash(existing_key)], (
f"map_existing_key must leave the user with only their pre-existing key, got {keys}"
)
mapping: Final = _mapping_for(client, identity.user_id)
assert mapping is not None, (
f"no JWT mapping row for sub={identity.user_id}: {unwrap(client.jwt_mapping_list())}"
)
assert mapping.jwt_claim_name == "sub", f"mapping must bind the sub claim, got {mapping}"
assert mapping.created_by == "auto_register", f"mapping must be written by auto_register, got {mapping}"
rows: Final = client.proxy.poll_logs_for_key(existing_key)
assert any(row.request_id == response.id for row in rows), (
f"the JWT chat must be billed to the user's existing key, spend rows for it: {rows}"
)
@pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless")
def test_first_jwt_call_mints_a_key_when_the_user_has_none(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
assert response.choices, f"JWT chat returned no completion: {response}"
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert len(keys) == 1, f"a keyless user must get exactly one minted key, got {keys}"
mapping: Final = _mapping_for(client, identity.user_id)
assert mapping is not None and mapping.jwt_claim_name == "sub", (
f"the minted key must be recorded as a sub-claim mapping, mappings: {unwrap(client.jwt_mapping_list())}"
)
@pytest.mark.covers("other.auth.jwt.auto_register_default_mints")
def test_default_behavior_still_mints_when_the_user_already_has_a_key(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
existing_key: Final = client.proxy.generate_key(
KeyGenerateBody(user_id=identity.user_id, key_alias=f"e2e-jwt-existing-{unique_marker()}")
)
resources.defer(lambda: client.proxy.delete_key(existing_key))
response: Final = unwrap(minting_gateway.proxy.chat(idp.access_token(identity), _ping()))
assert response.id is not None, f"JWT chat returned no response id: {response}"
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert len(keys) == 2, (
f"default auto_register must mint a second key for a user who already has one, got {keys}"
)
rows: Final = client.proxy.poll_logs_for_request_id(response.id)
assert rows and all(row.api_key != _key_hash(existing_key) for row in rows), (
f"the default path must bill the minted key, not the user's existing one: {rows}"
)

View file

@ -16,6 +16,7 @@ markers =
quiet_stack: measures the proxy itself, so it runs while no other test on this host is hitting the stack; every other test waits for it to finish
mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set
provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set
owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL on the pytest host; deselected unless E2E_OWNED_GATEWAY is set
otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set
otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set
secret_manager: needs a proxy booted from gateway/secret_manager_<system>_ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)

View file

@ -88,15 +88,89 @@ class DatabaseRelay:
)
class HeldStatementRelay:
def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None:
self.port: Final = _free_port()
self._upstream_host: Final = upstream_host
self._upstream_port: Final = upstream_port
self._trigger: Final = trigger
self._loop: Final = asyncio.new_event_loop()
self._released: Final = asyncio.Event()
self.held: Final = threading.Event()
self._ready: Final = threading.Event()
self._thread: Final = threading.Thread(target=self._run, daemon=True)
def release(self) -> None:
self._loop.call_soon_threadsafe(self._released.set)
def start(self) -> None:
self._thread.start()
assert self._ready.wait(10), "Database relay did not start"
def stop(self) -> None:
self.release()
self._loop.call_soon_threadsafe(self._loop.stop)
self._thread.join(10)
def _run(self) -> None:
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port))
self._ready.set()
self._loop.run_forever()
def _holds(self, window: bytes) -> bool:
return not self.held.is_set() and self._trigger in window
async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None:
server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port)
async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None:
tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches
try:
while chunk := await reader.read(65536):
window: Final = tail + chunk
if inspect and self._holds(window):
self.held.set()
await self._released.wait()
tail = window[-(len(self._trigger) - 1) :]
writer.write(chunk)
await writer.drain()
except (ConnectionError, asyncio.IncompleteReadError):
return
finally:
writer.close()
await asyncio.gather(
forward(client_reader, server_writer, True),
forward(server_reader, client_writer, False),
)
def _relayed_url(database_url: str, port: int) -> str:
parts: Final = urlsplit(database_url)
credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else ""
return urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{port}"))
@contextmanager
def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger)
relay.start()
credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else ""
relayed: Final = urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{relay.port}"))
try:
yield relay, relayed
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()
@contextmanager
def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[HeldStatementRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = HeldStatementRelay(parts.hostname, parts.port, trigger)
relay.start()
try:
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()

View file

@ -0,0 +1,219 @@
import json
import os
import time
import uuid
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import jwt
import pytest
import yaml
from cryptography.hazmat.primitives.asymmetric import rsa
from tests.integration._support.client import Gateway, eventually, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.database_relay import held_statement_relay
from tests.integration._support.process import owned_proxy
from tests.integration._support.wire import Reply, Request, wire_server
KEY_ID: Final = "integration-jwt-map-existing-key"
MAPPING_INSERT: Final = b'INSERT INTO "public"."LiteLLM_JWTKeyMapping"'
pytestmark = pytest.mark.timeout(240)
def _hash(key: str) -> str:
return sha256(key.encode()).hexdigest()
def _config(directory: Path, claim_field: str) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"] = {
**config["general_settings"],
"enable_jwt_auth": True,
"litellm_jwtauth": {
"user_id_jwt_field": "sub",
"user_email_jwt_field": "email",
"virtual_key_claim_field": claim_field,
"unregistered_jwt_client_behavior": "auto_register",
"auto_register_map_existing_key": True,
},
}
path: Final = directory / f"jwt_map_existing_key_{claim_field}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@contextmanager
def _issuer() -> Iterator[tuple[rsa.RSAPrivateKey, str]]:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_jwk: Final = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key()))
jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return Reply(body=jwks)
with wire_server(respond) as server:
yield private_key, server.url
def _token(private_key: rsa.RSAPrivateKey, subject: str, **claims: str) -> str:
now: Final = int(time.time())
return jwt.encode(
{"sub": subject, **claims, "iat": now, "exp": now + 300},
private_key,
algorithm="RS256",
headers={"kid": KEY_ID},
)
def _chat(candidate: Gateway, model: str, token: str) -> httpx.Response:
return candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "map existing key control"}]},
key=token,
)
def _mapped_token(claim_name: str, claim_value: str) -> str:
rows: Final = read_rows(
'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s',
(claim_name, claim_value),
)
assert len(rows) == 1, rows
return string_value(rows[0]["token"])
def _user_key_hashes(user: str) -> frozenset[str]:
rows: Final = read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,))
return frozenset(string_value(row["token"]) for row in rows)
def _billed_key(response: httpx.Response) -> str:
rows: Final = eventually(
lambda: read_rows(
'SELECT api_key FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (str(response.json()["id"]),)
),
lambda values: len(values) == 1,
seconds=70,
)
return string_value(rows[0]["api_key"])
def test_first_jwt_call_reuses_the_newest_durable_llm_key_and_skips_every_ineligible_newer_key(
gateway: Gateway, tmp_path: Path
) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user(user_role="internal_user")
older_durable: Final = scenario.key(user_id=user)
durable: Final = scenario.key(user_id=user)
skipped: Final = {
"older_durable": older_durable,
"expiring": scenario.key(user_id=user, duration="1h"),
"management_only": scenario.key(user_id=user, allowed_routes=["management_routes"]),
"auto_registered_look_alike": scenario.key(user_id=user, metadata={"auto_registered": True}),
"other_team": scenario.key(user_id=user, team_id=scenario.team()),
"blocked": scenario.key(user_id=user),
}
gateway.post("/key/block", {"key": skipped["blocked"]})
keys_before: Final = _user_key_hashes(user)
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub")
) as candidate:
response: Final = _chat(candidate, model, _token(private_key, user))
assert response.status_code == 200, response.text
mapped: Final = _mapped_token("sub", user)
assert mapped == _hash(durable), {
"mapped_to": next((name for name, key in skipped.items() if _hash(key) == mapped), mapped)
}
assert _user_key_hashes(user) == keys_before, "a key was minted although a reusable one existed"
assert _billed_key(response) == _hash(durable)
def test_user_matched_by_email_instead_of_sub_still_reuses_their_existing_key(gateway: Gateway, tmp_path: Path) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
email: Final = f"integration-{uuid.uuid4().hex}@example.com"
user: Final = scenario.user(user_role="internal_user", user_email=email)
existing: Final = scenario.key(user_id=user)
subject: Final = f"integration-idp-subject-{uuid.uuid4().hex}"
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub")
) as candidate:
response: Final = _chat(candidate, model, _token(private_key, subject, email=email.upper()))
assert response.status_code == 200, response.text
assert _mapped_token("sub", subject) == _hash(existing)
assert _user_key_hashes(user) == frozenset({_hash(existing)}), "a key was minted for an email-matched user"
assert _billed_key(response) == _hash(existing)
def test_shared_client_claim_never_maps_a_second_user_onto_the_first_users_personal_key(
gateway: Gateway, tmp_path: Path
) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
first_user: Final = scenario.user(user_role="internal_user")
second_user: Final = scenario.user(user_role="internal_user")
personal: Final = scenario.key(user_id=first_user)
client_id: Final = f"integration-shared-client-{uuid.uuid4().hex}"
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "client_id")
) as candidate:
first: Final = _chat(candidate, model, _token(private_key, first_user, client_id=client_id))
second: Final = _chat(candidate, model, _token(private_key, second_user, client_id=client_id))
assert first.status_code == 200, first.text
assert second.status_code == 200, second.text
mapped: Final = _mapped_token("client_id", client_id)
assert mapped != _hash(personal), "the shared client claim was mapped to the first user's personal key"
assert (_billed_key(first), _billed_key(second)) == (mapped, mapped)
assert read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s', (_hash(personal),)) == []
def test_concurrent_first_jwt_calls_of_a_keyless_user_both_succeed_on_one_surviving_mapped_key(
gateway: Gateway, tmp_path: Path
) -> None:
writer_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"]
with (
_issuer() as (private_key, jwks_url),
gateway.scenario() as scenario,
held_statement_relay(writer_url, MAPPING_INSERT) as (relay, relayed_url),
):
model: Final = scenario.model()
user: Final = scenario.user(user_role="internal_user")
token: Final = _token(private_key, user)
overrides: Final = {
"JWT_PUBLIC_KEY_URL": jwks_url,
"DATABASE_URL": relayed_url,
"PRISMA_HEALTH_WATCHDOG_ENABLED": "false",
}
with (
owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, "sub")) as candidate,
ThreadPoolExecutor(max_workers=1) as pool,
):
held_call: Final = pool.submit(_chat, candidate, model, token)
assert relay.held.wait(60), "the first call never reached its mapping insert"
racing: Final = _chat(candidate, model, token)
relay.release()
held: Final = held_call.result(timeout=60)
assert racing.status_code == 200, racing.text
assert held.status_code == 200, held.text
keys: Final = _user_key_hashes(user)
assert len(keys) == 1, keys
assert _mapped_token("sub", user) in keys
assert (_billed_key(held), _billed_key(racing)) == (_mapped_token("sub", user),) * 2

View file

@ -2020,6 +2020,454 @@ async def test_auto_register_binds_api_key_to_token_hash():
assert result.end_user_id == "validated-end-user"
def _auto_register_patches(*, plaintext_key: str | None = "sk-minted-plaintext"):
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import CredentialRef
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.proxy_server import hash_token
resolved_key = UserAPIKeyAuth(
token="existing-hash" if plaintext_key is None else hash_token(plaintext_key),
user_id="validated-user",
team_id="validated-team",
org_id="key-own-org",
)
principal = IdentityStore._principal_from_key(
resolved_key,
auth_method=AuthMethod.API_KEY,
credential_ref=CredentialRef(token_id=resolved_key.token),
)
return (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": plaintext_key},
),
patch(
"litellm.proxy.auth.resolvers.store.IdentityStore.resolve",
new_callable=AsyncMock,
return_value=principal,
),
)
def _auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, **over):
kwargs = {
"virtual_key_claim_field": "sub",
"claim_value": "validated-user",
"jwt_handler": jwt_handler,
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"parent_otel_span": None,
"proxy_logging_obj": MagicMock(),
"cache_key": "jwt_key_mapping:sub:validated-user",
"team_id": "validated-team",
"user_id": "validated-user",
"org_id": "jwt-org",
"end_user_id": "validated-end-user",
}
kwargs.update(over)
return kwargs
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_reuses_users_key_but_never_an_auto_registered_one():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[
{"token": "auto-registered-hash", "metadata": {"auto_registered": True}},
{"token": "existing-hash", "metadata": {}},
]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with generate_patch as generate_key, resolve_patch:
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
generate_key.assert_not_awaited()
create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]
assert create_data["token"] == "existing-hash"
assert create_data["created_by"] == "auto_register"
assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "existing-hash"
assert result is not None
assert result.token == "existing-hash"
assert result.api_key == "existing-hash"
assert result.org_id == "key-own-org"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_mints_when_user_has_no_key():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
generate_key.assert_awaited_once()
create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]
assert create_data["token"] == hash_token("sk-minted-plaintext")
assert result is not None
assert result.token == hash_token("sk-minted-plaintext")
@pytest.mark.asyncio
async def test_auto_register_default_never_looks_up_existing_keys():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub", virtual_key_mapping_cache_ttl=300)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler))
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_race_loser_keeps_reused_key():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(side_effect=Exception("Unique constraint failed (P2002)"))
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with (
generate_patch,
resolve_patch,
patch(
"litellm.proxy.auth.user_api_key_auth.get_jwt_key_mapping_object",
new_callable=AsyncMock,
return_value="winner-hash",
),
):
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
assert result is not None
assert result.org_id == "key-own-org"
prisma_client.db.litellm_verificationtoken.delete.assert_not_awaited()
assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "winner-hash"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_user_id_none_mints():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, user_id=None)
)
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_reuses_when_the_user_was_matched_by_a_fallback_lookup():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
user_email_jwt_field="email",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client,
user_api_key_cache,
jwt_handler,
claim_value="idp-subject-not-the-db-user-id",
cache_key="jwt_key_mapping:sub:idp-subject-not-the-db-user-id",
)
)
generate_key.assert_not_awaited()
assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == "existing-hash"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_mints_when_the_claim_is_not_a_user_identity_field():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client,
user_api_key_cache,
jwt_handler,
virtual_key_claim_field="azp",
claim_value="shared-client-app",
cache_key="jwt_key_mapping:azp:shared-client-app",
)
)
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == hash_token(
"sk-minted-plaintext"
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("issuer_user_id_field", "expect_reuse"),
[("uid", False), (None, True)],
)
async def test_auto_register_map_existing_key_uses_the_issuers_own_user_field_over_the_global_one(
issuer_user_id_field, expect_reuse
):
from litellm.proxy._types import JWTIssuerConfig
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
issuers=[
JWTIssuerConfig(
issuer="https://idp.example.com", audience="litellm", user_id_jwt_field=issuer_user_id_field
)
],
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client, user_api_key_cache, jwt_handler, jwt_issuer="https://idp.example.com"
)
)
mapped_token = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"]
assert mapped_token == ("existing-hash" if expect_reuse else hash_token("sk-minted-plaintext"))
assert generate_key.await_count == (0 if expect_reuse else 1)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("map_existing_key", "master_key", "reused_key_models", "expect_denied"),
[
(True, "sk-master", ["some-other-model"], True),
(True, "sk-master", [], False),
(False, "sk-master", ["some-other-model"], False),
(True, None, ["some-other-model"], False),
],
)
async def test_auto_register_map_existing_key_first_request_runs_key_checks(
map_existing_key: bool, master_key: str | None, reused_key_models: list[str], expect_denied: bool
) -> None:
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
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_register_map_existing_key=map_existing_key,
)
reused_key = UserAPIKeyAuth(
token="hashed-existing-key",
api_key="hashed-existing-key",
user_id="validated-user",
team_id="validated-team",
models=reused_key_models,
)
mock_jwt_result = {
"is_proxy_admin": False,
"team_object": None,
"user_object": LiteLLM_UserTable(user_id="validated-user", user_role="internal_user"),
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": "validated-team",
"user_id": "validated-user",
"user_email": None,
"end_user_id": None,
"org_id": None,
"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", {"enable_jwt_auth": True}),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", master_key),
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(post_call_failure_hook=AsyncMock(return_value=None)),
),
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=reused_key,
),
):
call = _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"},
)
if expect_denied:
with pytest.raises(ProxyException, match="not available for this API key"):
await call
return
result = await call
assert result.api_key == "hashed-existing-key"
assert result.user_id == "validated-user"
assert result.team_id == "validated-team"
assert result.models == reused_key_models
@pytest.mark.asyncio
@pytest.mark.parametrize("active", [True, False])
async def test_auto_register_first_request_propagates_user_email(active: bool) -> None: