diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index df950cb6a73..351cce75fdf 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -12,7 +12,7 @@ import fnmatch import re import secrets from datetime import datetime, timezone -from typing import Any, Iterator, List, Optional, Tuple, Union, cast +from typing import Any, Iterator, List, NamedTuple, Optional, Tuple, Union, cast import fastapi from fastapi import HTTPException, Request, WebSocket, status @@ -587,38 +587,22 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints( return api_key -def _resolve_inherited_team_id( - jwt_handler: JWTHandler, - jwt_claims: dict, -) -> Optional[str]: +class _PendingAutoRegister(NamedTuple): """ - Resolve the team_id the JWT would have been bound to under standard team- - based auth. AUTO_REGISTER stamps this on the new key so it inherits the - team's budget/model/tpm/rpm/route limits — auto-registered clients then - have the same constraints they'd have via the JWT team-auth path. + Signal returned by ``_resolve_jwt_to_virtual_key`` when the JWT's claim is + unmapped and ``unregistered_jwt_client_behavior`` is AUTO_REGISTER. - If ``team_id_jwt_field`` is configured but the JWT lacks the claim and no - ``team_id_default`` is set, raise 403: the operator has declared every JWT - must carry a team, and registering a teamless (unbounded) key here would - silently grant broader access than the team-auth path allows. + The caller MUST run standard ``JWTAuthManager.auth_builder`` to apply RBAC, + scope mappings, ``custom_validate``, and ``user_allowed_email_domain`` + policy BEFORE calling ``_auto_register_jwt_mapping`` with the validated + ``team_id`` / ``user_id`` from the auth_builder result. Auto-registering + purely on a signature-valid JWT (the old behavior) bypassed every JWT + policy beyond signature verification. """ - if ( - jwt_handler.litellm_jwtauth.team_id_jwt_field is None - and jwt_handler.litellm_jwtauth.team_id_default is None - ): - return None - team_id = jwt_handler.get_team_id(token=jwt_claims, default_value=None) - if team_id is None: - raise HTTPException( - status_code=403, - detail=( - "JWT Key Mapping: AUTO_REGISTER requires a team_id on the JWT " - "(via team_id_jwt_field) or a configured team_id_default — " - "refusing to create an unbounded key. Access denied." - ), - ) - return team_id + claim_field: str + claim_value: str + cache_key: str async def _auto_register_jwt_mapping( @@ -631,13 +615,15 @@ async def _auto_register_jwt_mapping( proxy_logging_obj: ProxyLogging, cache_key: str, team_id: Optional[str] = None, + user_id: Optional[str] = None, ) -> Optional[UserAPIKeyAuth]: """ Auto-register: create a new virtual key + mapping for an unrecognised JWT - claim value. The key is stamped with the JWT's resolved ``team_id`` (when - available) so it inherits the team's budget/model/tpm/rpm/route limits via - the standard auth path — auto-registered clients are then no more permissive - than the team-based JWT auth path they would otherwise have taken. + claim value. ``team_id`` and ``user_id`` must come from a successful + ``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER + RBAC/scope/custom_validate/email-domain policy has been enforced. The key + is stamped with those values so the cached future-request path inherits + the same team/user limits the auth_builder path would have applied. Race safety: if two concurrent requests both reach here simultaneously (both saw no mapping in the DB), one will win the unique-constraint race on @@ -659,6 +645,7 @@ async def _auto_register_jwt_mapping( request_type="key", table_name="key", team_id=team_id, + user_id=user_id, metadata={ "auto_registered": True, "jwt_claim_field": virtual_key_claim_field, @@ -754,7 +741,22 @@ async def _resolve_jwt_to_virtual_key( user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, -) -> Optional[UserAPIKeyAuth]: +) -> Union[Optional[UserAPIKeyAuth], "_PendingAutoRegister"]: + """ + Returns: + - ``UserAPIKeyAuth``: a resolved virtual key (cache hit or DB hit). The + caller may use this directly; JWT policy has been enforced previously + (at key-creation time or, for cached results, before caching). + - ``_PendingAutoRegister``: claim is unmapped and behavior is AUTO_REGISTER. + The caller MUST run ``JWTAuthManager.auth_builder`` to enforce JWT + policy (RBAC, scope, custom_validate, email-domain), then invoke + ``_auto_register_jwt_mapping`` with the validated team_id/user_id. + - ``None``: claim is unmapped and behavior is FALLBACK_TEAM_MAPPING. + The caller falls through to standard team-based JWT auth (which itself + enforces full JWT policy via auth_builder). + - Raises HTTPException: REJECT policy hit, missing claim under + REJECT/AUTO_REGISTER, or other policy violations. + """ virtual_key_claim_field = jwt_handler.litellm_jwtauth.virtual_key_claim_field if virtual_key_claim_field is None: return None @@ -800,7 +802,7 @@ async def _resolve_jwt_to_virtual_key( ) if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER: # Stale sentinel written under a prior fallback_team_mapping config — - # evict it and auto-register now that the policy has changed. Raise + # evict it and defer auto-register to after auth_builder runs. Raise # the same 500 as the fresh-path AUTO_REGISTER branch when there is # no DB, so behavior is consistent regardless of whether the cache # happens to hold the sentinel. @@ -812,18 +814,11 @@ async def _resolve_jwt_to_virtual_key( "Configure a database or change unregistered_jwt_client_behavior." ), ) - inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims) await user_api_key_cache.async_delete_cache(cache_key) - return await _auto_register_jwt_mapping( - virtual_key_claim_field=virtual_key_claim_field, + return _PendingAutoRegister( + claim_field=virtual_key_claim_field, claim_value=str(claim_value), - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, cache_key=cache_key, - team_id=inherited_team_id, ) return None elif cached_mapping is not None: @@ -884,17 +879,14 @@ async def _resolve_jwt_to_virtual_key( "Configure a database or change unregistered_jwt_client_behavior." ), ) - inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims) - return await _auto_register_jwt_mapping( - virtual_key_claim_field=virtual_key_claim_field, + # Defer: caller runs JWTAuthManager.auth_builder to enforce RBAC, scope, + # custom_validate, and email-domain policy, then auto-registers using + # the validated identity. Auto-registering here on a signature-only + # JWT would bypass every JWT policy beyond signature verification. + return _PendingAutoRegister( + claim_field=virtual_key_claim_field, claim_value=str(claim_value), - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, cache_key=cache_key, - team_id=inherited_team_id, ) # FALLBACK_TEAM_MAPPING (default): cache the miss and return None so the @@ -1093,6 +1085,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Try JWT-to-Virtual-Key mapping first to avoid # unnecessary DB queries in auth_builder do_standard_jwt_auth = True + pending_auto_register: Optional[_PendingAutoRegister] = None if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None: # Decode JWT to get claims without running full auth_builder jwt_claims: Optional[dict] @@ -1101,7 +1094,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 else: jwt_claims = await jwt_handler.auth_jwt(token=api_key) - valid_token = await _resolve_jwt_to_virtual_key( + resolve_result = await _resolve_jwt_to_virtual_key( jwt_claims=jwt_claims, jwt_handler=jwt_handler, prisma_client=prisma_client, @@ -1109,11 +1102,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - if valid_token is not None: + if isinstance(resolve_result, UserAPIKeyAuth): + valid_token = resolve_result api_key = valid_token.token or "" valid_token.jwt_claims = jwt_claims do_standard_jwt_auth = False # Fall through to virtual key checks + elif isinstance(resolve_result, _PendingAutoRegister): + # Run full JWT policy (RBAC, scope, custom_validate, + # email-domain) via auth_builder, then create the key + # from the validated identity below. + pending_auto_register = resolve_result + # else: None → FALLBACK_TEAM_MAPPING, falls through to + # standard JWT auth_builder below if do_standard_jwt_auth: with tracer.trace("litellm.proxy.auth.jwt_auth_builder"): @@ -1229,6 +1230,30 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 else None ) + # AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key. + # JWT policy (RBAC, scope, custom_validate, email-domain) + # has now been enforced by auth_builder above. Create the + # mapping + virtual key from the *validated* identity, then + # replace valid_token with the new key so downstream checks + # use the key-scoped path. + if pending_auto_register is not None and prisma_client is not None: + auto_registered = await _auto_register_jwt_mapping( + virtual_key_claim_field=pending_auto_register.claim_field, + claim_value=pending_auto_register.claim_value, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + cache_key=pending_auto_register.cache_key, + team_id=team_id, + user_id=user_id, + ) + if auto_registered is not None: + auto_registered.jwt_claims = jwt_claims + valid_token = auto_registered + api_key = valid_token.token or "" + # Check if model has zero cost - if so, skip all budget checks model = _get_model_from_request_context( request_data=request_data, diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index dc47f55a314..15c3834d1a6 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -596,14 +596,17 @@ async def test_reject_behavior_raises_403_on_cached_no_mapping(): @pytest.mark.asyncio -async def test_auto_register_creates_key_and_mapping(): +async def test_auto_register_returns_pending_signal_without_creating_key(): """ - When unregistered_jwt_client_behavior='auto_register' and no mapping exists, - _resolve_jwt_to_virtual_key must create a key + mapping row and return a - UserAPIKeyAuth object. The mapping row stores the hashed token (FK to - LiteLLM_VerificationToken), not the plaintext key. + Security: when unregistered_jwt_client_behavior='auto_register' and no + mapping exists, _resolve_jwt_to_virtual_key must NOT create the key yet. + It returns a _PendingAutoRegister signal so the caller can run + JWTAuthManager.auth_builder (enforcing RBAC, scope mappings, + custom_validate, user_allowed_email_domain) FIRST. Creating the key here + would bypass every JWT policy beyond signature verification. """ - from litellm.proxy._types import UnregisteredJWTClientBehavior, hash_token + from litellm.proxy._types import UnregisteredJWTClientBehavior + from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( @@ -617,10 +620,55 @@ async def test_auto_register_creates_key_and_mapping(): prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + user_api_key_cache = DualCache() + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_gen_key: + result = await _resolve_jwt_to_virtual_key( + jwt_claims=jwt_claims, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert isinstance(result, _PendingAutoRegister) + assert result.claim_field == "sub" + assert result.claim_value == "new-user-42" + assert result.cache_key == "jwt_key_mapping:sub:new-user-42" + # CRITICAL: no key was created — that must wait until after auth_builder + mock_gen_key.assert_not_called() + prisma_client.db.litellm_jwtkeymapping.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_auto_register_creates_key_and_mapping_when_helper_invoked(): + """ + When the caller invokes _auto_register_jwt_mapping directly (after + auth_builder validation), the helper creates the key + mapping row and + returns a UserAPIKeyAuth. The mapping row stores the hashed token (FK to + LiteLLM_VerificationToken), not the plaintext key. + """ + from litellm.proxy._types import hash_token + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + virtual_key_mapping_cache_ttl=300, + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + user_api_key_cache = DualCache() plaintext_key = "sk-auto-key" expected_hash = hash_token(plaintext_key) - mock_key_obj = UserAPIKeyAuth(token=expected_hash, team_id=None) + mock_key_obj = UserAPIKeyAuth(token=expected_hash, team_id="validated-team") with ( patch( @@ -635,35 +683,45 @@ async def test_auto_register_creates_key_and_mapping(): mock_gen_key.return_value = {"token": plaintext_key, "key": plaintext_key} mock_get_key.return_value = mock_key_obj - result = await _resolve_jwt_to_virtual_key( - jwt_claims=jwt_claims, + result = await _auto_register_jwt_mapping( + virtual_key_claim_field="sub", + claim_value="new-user-42", jwt_handler=jwt_handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, proxy_logging_obj=None, + cache_key="jwt_key_mapping:sub:new-user-42", + team_id="validated-team", + user_id="validated-user", ) assert result == mock_key_obj - # Mapping row must have been created with the hashed token (FK target) - prisma_client.db.litellm_jwtkeymapping.create.assert_called_once() + # generate_key_helper_fn was passed table_name="key" (not user-upsert path) + # and the validated team_id + user_id from auth_builder + assert mock_gen_key.call_args.kwargs["table_name"] == "key" + assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team" + assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user" + # Mapping row was created with the hashed token (FK target) call_data = prisma_client.db.litellm_jwtkeymapping.create.call_args[1]["data"] assert call_data["jwt_claim_name"] == "sub" assert call_data["jwt_claim_value"] == "new-user-42" assert call_data["token"] == expected_hash - # Cache must hold the hashed token cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:new-user-42") assert cached == expected_hash @pytest.mark.asyncio -async def test_auto_register_triggers_on_stale_no_mapping_sentinel(): +async def test_auto_register_returns_pending_signal_on_stale_no_mapping_sentinel(): """ If the cache holds a stale __NO_MAPPING__ sentinel (written under a prior - fallback_team_mapping config) and behavior is now AUTO_REGISTER, the sentinel - must be evicted and auto-registration must run — not silently return None. + fallback_team_mapping config) and behavior is now AUTO_REGISTER, the + resolver must evict the sentinel and return _PendingAutoRegister (so the + caller can run auth_builder before creating the key) — not silently return + None and not create the key on the spot. """ from litellm.proxy._types import UnregisteredJWTClientBehavior + from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( @@ -678,26 +736,14 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel(): prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() user_api_key_cache = DualCache() - # Seed the stale sentinel await user_api_key_cache.async_set_cache( "jwt_key_mapping:email:alice@corp.com", "__NO_MAPPING__" ) - mock_key_obj = UserAPIKeyAuth(token="hashed_auto_key", team_id=None) - - with ( - patch( - "litellm.proxy.auth.user_api_key_auth.get_key_object", - new_callable=AsyncMock, - ) as mock_get_key, - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_gen_key, - ): - mock_gen_key.return_value = {"token": "hashed_auto_key", "key": "sk-auto-key"} - mock_get_key.return_value = mock_key_obj - + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_gen_key: result = await _resolve_jwt_to_virtual_key( jwt_claims=jwt_claims, jwt_handler=jwt_handler, @@ -707,9 +753,15 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel(): proxy_logging_obj=None, ) - # Must have auto-registered, not returned None - assert result == mock_key_obj - prisma_client.db.litellm_jwtkeymapping.create.assert_called_once() + assert isinstance(result, _PendingAutoRegister) + # Stale sentinel must be evicted so the deferred auto-register actually + # runs after auth_builder validates the JWT + cached_after = await user_api_key_cache.async_get_cache( + "jwt_key_mapping:email:alice@corp.com" + ) + assert cached_after is None + mock_gen_key.assert_not_called() + prisma_client.db.litellm_jwtkeymapping.create.assert_not_called() @pytest.mark.asyncio @@ -1123,33 +1175,33 @@ async def test_auto_register_raises_503_when_winner_mapping_vanishes(): # ────────────────────────────────────────────── -# Tests: AUTO_REGISTER inherits team_id from JWT +# Tests: AUTO_REGISTER stamps validated identity from auth_builder # ────────────────────────────────────────────── @pytest.mark.asyncio -async def test_auto_register_stamps_team_id_from_jwt_claim(): +async def test_auto_register_helper_stamps_validated_team_and_user(): """ - When team_id_jwt_field is configured and the JWT carries the claim, - AUTO_REGISTER must stamp the resolved team_id on the new key so it inherits - the team's limits via standard auth (budget, models, tpm/rpm, allowed_routes). + The deferred-auto-register contract: _auto_register_jwt_mapping is called + with team_id and user_id from JWTAuthManager.auth_builder's *validated* + result (after RBAC, scope mappings, custom_validate, email-domain policy). + These must be passed to generate_key_helper_fn so the created key carries + them — the cached future-request path then inherits the same team/user + limits the auth_builder path would have applied. """ - from litellm.proxy._types import UnregisteredJWTClientBehavior + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( virtual_key_claim_field="sub", - team_id_jwt_field="team_id", - unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER, + virtual_key_mapping_cache_ttl=300, ) - jwt_claims = {"sub": "new-user", "team_id": "team-engineering"} prisma_client = MagicMock() - prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() - - user_api_key_cache = DualCache() - mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="team-engineering") + mock_key_obj = UserAPIKeyAuth( + token="hashed", team_id="validated-team", user_id="validated-user" + ) with ( patch( @@ -1164,163 +1216,23 @@ async def test_auto_register_stamps_team_id_from_jwt_claim(): mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"} mock_get_key.return_value = mock_key_obj - result = await _resolve_jwt_to_virtual_key( - jwt_claims=jwt_claims, + result = await _auto_register_jwt_mapping( + virtual_key_claim_field="sub", + claim_value="new-user", jwt_handler=jwt_handler, prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, + user_api_key_cache=DualCache(), parent_otel_span=None, proxy_logging_obj=None, + cache_key="jwt_key_mapping:sub:new-user", + team_id="validated-team", + user_id="validated-user", ) assert result == mock_key_obj - # generate_key_helper_fn must have been called with team_id from the JWT, - # so the key inherits team-level limits via the standard auth path. - mock_gen_key.assert_called_once() - assert mock_gen_key.call_args.kwargs["team_id"] == "team-engineering" - - -@pytest.mark.asyncio -async def test_auto_register_rejects_when_team_id_required_but_missing_from_jwt(): - """ - Security: when team_id_jwt_field is configured (operator wants every JWT - bound to a team), AUTO_REGISTER must refuse a JWT that lacks the team - claim. Otherwise the auto-created key would be teamless and inherit no - budget/model/rate limits — broader access than the team-auth path allows. - """ - from litellm.proxy._types import UnregisteredJWTClientBehavior - - jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - virtual_key_claim_field="sub", - team_id_jwt_field="team_id", - unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER, - ) - jwt_claims = {"sub": "new-user"} # no "team_id" - - prisma_client = MagicMock() - prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) - - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_gen_key, - pytest.raises(HTTPException) as exc_info, - ): - await _resolve_jwt_to_virtual_key( - jwt_claims=jwt_claims, - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=DualCache(), - parent_otel_span=None, - proxy_logging_obj=None, - ) - - assert exc_info.value.status_code == 403 - assert "team_id" in exc_info.value.detail - # No key was created — we refused before generate_key_helper_fn ran - mock_gen_key.assert_not_called() - - -@pytest.mark.asyncio -async def test_auto_register_uses_team_id_default_when_jwt_lacks_team_claim(): - """ - When team_id_jwt_field is configured but the JWT lacks the claim and - team_id_default is set, the default team_id is stamped on the new key — - the explicit operator-chosen fallback team bounds the auto-registered key. - """ - from litellm.proxy._types import UnregisteredJWTClientBehavior - - jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - virtual_key_claim_field="sub", - team_id_jwt_field="team_id", - team_id_default="default-team", - unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER, - ) - jwt_claims = {"sub": "new-user"} # no "team_id" - - prisma_client = MagicMock() - prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) - prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() - - mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="default-team") - - with ( - patch( - "litellm.proxy.auth.user_api_key_auth.get_key_object", - new_callable=AsyncMock, - ) as mock_get_key, - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_gen_key, - ): - mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"} - mock_get_key.return_value = mock_key_obj - - result = await _resolve_jwt_to_virtual_key( - jwt_claims=jwt_claims, - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=DualCache(), - parent_otel_span=None, - proxy_logging_obj=None, - ) - - assert result == mock_key_obj - assert mock_gen_key.call_args.kwargs["team_id"] == "default-team" - - -@pytest.mark.asyncio -async def test_auto_register_no_team_id_when_team_field_not_configured(): - """ - When the operator has NOT configured team_id_jwt_field or team_id_default, - they have explicitly opted out of team-based bounding. AUTO_REGISTER then - creates a teamless key (preserving prior behavior) — there is no team to - inherit from. The key being unbounded matches what the team-auth path - would have produced in the same config. - """ - from litellm.proxy._types import UnregisteredJWTClientBehavior - - jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - virtual_key_claim_field="sub", - unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER, - ) - jwt_claims = {"sub": "new-user", "team_id": "ignored-because-not-configured"} - - prisma_client = MagicMock() - prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) - prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() - - mock_key_obj = UserAPIKeyAuth(token="hashed", team_id=None) - - with ( - patch( - "litellm.proxy.auth.user_api_key_auth.get_key_object", - new_callable=AsyncMock, - ) as mock_get_key, - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_gen_key, - ): - mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"} - mock_get_key.return_value = mock_key_obj - - await _resolve_jwt_to_virtual_key( - jwt_claims=jwt_claims, - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=DualCache(), - parent_otel_span=None, - proxy_logging_obj=None, - ) - - # team_id is None because the operator opted out of team-based config - assert mock_gen_key.call_args.kwargs["team_id"] is None + # Validated identity flowed to the key (key inherits team + user limits) + assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team" + assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user" # ──────────────────────────────────────────────