mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
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.
This commit is contained in:
parent
478de14bb5
commit
69b121c969
2 changed files with 37 additions and 16 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue