From 69b121c969589de000898fe9037ae94c1249fa51 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Wed, 23 Sep 2026 18:36:17 -0700 Subject: [PATCH] fix(jwt): keep the early return when no master key is set Without a master key the generic virtual-key path returns a bare INTERNAL_USER object, so falling through on the first auto-registered request dropped the key's team, models and budgets. Only fall through when a master key is configured. Tests now assert the reused key per team rather than the query shape, and cover the flag-off early return and the no-master-key case. --- litellm/proxy/auth/user_api_key_auth.py | 7 ++- .../proxy/auth/test_user_api_key_auth.py | 46 +++++++++++++------ 2 files changed, 37 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 7813578d385..6930376959b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1807,7 +1807,12 @@ async def _user_api_key_auth_builder( valid_token = auto_registered api_key = valid_token.token or "" - if auto_registered is None or not jwt_handler.litellm_jwtauth.auto_register_map_existing_key: + 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, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index a3b2badb5a5..75fc2ebceda 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2237,16 +2237,25 @@ async def test_auto_register_map_existing_key_skips_keys_that_cannot_call_llm_ro @pytest.mark.asyncio -@pytest.mark.parametrize("resolved_team_id", ["validated-team", None]) +@pytest.mark.parametrize( + ("resolved_team_id", "expected_token"), + [("validated-team", "team-key-hash"), (None, "teamless-key-hash")], +) async def test_auto_register_map_existing_key_only_reuses_keys_in_the_jwt_resolved_team( - resolved_team_id: str | None, + resolved_team_id: str | None, expected_token: str ) -> None: from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + keys_by_team: dict[str | None, str] = {"validated-team": "team-key-hash", None: "teamless-key-hash"} + + async def find_first(*, where: dict[str, object], order: dict[str, str]) -> SimpleNamespace | None: + if "team_id" not in where: + return SimpleNamespace(token="newest-key-in-any-team") + token = keys_by_team.get(where["team_id"]) + return None if token is None else SimpleNamespace(token=token) + prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( - return_value=SimpleNamespace(token="existing-hash") - ) + prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(side_effect=find_first) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() user_api_key_cache = MagicMock() @@ -2259,14 +2268,13 @@ async def test_auto_register_map_existing_key_only_reuses_keys_in_the_jwt_resolv ) generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) - with generate_patch, resolve_patch: + with generate_patch as generate_key, resolve_patch: await _auto_register_jwt_mapping( **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, team_id=resolved_team_id) ) - where = prisma_client.db.litellm_verificationtoken.find_first.await_args.kwargs["where"] - assert "team_id" in where, f"reuse must be scoped to the JWT-resolved team: {where}" - assert where["team_id"] == resolved_team_id + generate_key.assert_not_awaited() + assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == expected_token @pytest.mark.asyncio @@ -2399,11 +2407,16 @@ async def test_auto_register_map_existing_key_user_id_none_mints(): @pytest.mark.asyncio @pytest.mark.parametrize( - ("reused_key_models", "expect_denied"), - [(["some-other-model"], True), ([], False)], + ("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( - reused_key_models: list[str], expect_denied: bool + 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() @@ -2414,12 +2427,13 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( virtual_key_claim_field="sub", virtual_key_mapping_cache_ttl=300, - auto_register_map_existing_key=True, + 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 = { @@ -2429,7 +2443,7 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( "end_user_object": None, "org_object": None, "token": jwt_token, - "team_id": None, + "team_id": "validated-team", "user_id": "validated-user", "user_email": None, "end_user_id": None, @@ -2448,7 +2462,7 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( 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", "sk-master"), + 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( @@ -2493,6 +2507,8 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( 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