diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml index 8e03a902383..228e23f60d7 100644 --- a/.github/workflows/test-e2e-changed.yml +++ b/.github/workflows/test-e2e-changed.yml @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ef4c545507b..cabb04cc52b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 3400dccf2a7..e82f3eed7cc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 ## diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index d02c2114136..b20fb47306e 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -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}) diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 62153e38a83..37f0bf00da6 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -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", diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 0b9249d7420..39747607531 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -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"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 3fa9f534ffd..e88bfad8388 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -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. diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index b328c81687b..029b0135900 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -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 diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 83d8a9884a3..e027c410e44 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -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] = [] diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py index 93c198586f6..d7bad4f1ed1 100644 --- a/tests/e2e/other/other_client.py +++ b/tests/e2e/other/other_client.py @@ -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( diff --git a/tests/e2e/other/owned_jwt_gateway.py b/tests/e2e/other/owned_jwt_gateway.py new file mode 100644 index 00000000000..1af348cac60 --- /dev/null +++ b/tests/e2e/other/owned_jwt_gateway.py @@ -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 diff --git a/tests/e2e/other/test_jwt_auto_register_e2e.py b/tests/e2e/other/test_jwt_auto_register_e2e.py new file mode 100644 index 00000000000..8f7c2c6a693 --- /dev/null +++ b/tests/e2e/other/test_jwt_auto_register_e2e.py @@ -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}" + ) diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index e795ebe5721..dbd3ff47daa 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -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__ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py) diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py index 1b3bc0183a5..cb8b9098a64 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -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() diff --git a/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py new file mode 100644 index 00000000000..5215054a364 --- /dev/null +++ b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py @@ -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 diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 781d0a13bfd..a7219ac059b 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -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: