mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
fix(jwt): only reuse keys that can call LLM routes in auto_register_map_existing_key
Skip Admin UI session keys and keys whose allowed_routes restrict them to anything other than llm_api_routes (management, read_only, password-reset sessions). Mapping a JWT to one of those left the user with 401s or 403s on every LLM call, since the mapping persists.
This commit is contained in:
parent
50575e32ed
commit
449c176ef8
4 changed files with 55 additions and 5 deletions
|
|
@ -5303,8 +5303,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
description=(
|
||||
"Only used with unregistered_jwt_client_behavior='auto_register'. When True, the JWT claim is "
|
||||
"mapped to a virtual key the resolved internal user already owns instead of minting a new one. "
|
||||
"If the user owns several, the most recently created non-expired key is chosen. A new key is "
|
||||
"only minted when the user has none."
|
||||
"If the user owns several, the most recently created key that can call LLM routes is chosen: "
|
||||
"not blocked, not expired, not an Admin UI session key, and with no route restriction other than "
|
||||
"llm_api_routes. A new key is only minted when the user has no such key."
|
||||
),
|
||||
)
|
||||
routing_overrides: list[JWTRoutingOverride] | None = Field(
|
||||
|
|
|
|||
|
|
@ -931,7 +931,7 @@ class _PendingAutoRegister(NamedTuple):
|
|||
|
||||
|
||||
async def _latest_active_key_hash_for_user(prisma_client: PrismaClient, user_id: str) -> str | None:
|
||||
row: Final = await VerificationTokenRepository(prisma_client).find_latest_active_row_by_user_id(user_id)
|
||||
row: Final = await VerificationTokenRepository(prisma_client).find_latest_llm_api_row_by_user_id(user_id)
|
||||
return None if row is None else row.token
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from datetime import datetime, timezone
|
|||
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,8 +124,7 @@ 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_latest_active_row_by_user_id(self, user_id: str) -> "PrismaVerificationToken | None":
|
||||
"""Find the most recently created non-blocked, non-expired token row for a user."""
|
||||
async def find_latest_llm_api_row_by_user_id(self, user_id: str) -> "PrismaVerificationToken | None":
|
||||
row: Final = await self.table.find_first(
|
||||
where={ # mutable-ok: the prisma where clause contract is a plain dict
|
||||
"user_id": user_id,
|
||||
|
|
@ -136,6 +136,18 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
{"expires": {"gt": datetime.now(timezone.utc)}}, # mutable-ok: prisma filter literal
|
||||
]
|
||||
},
|
||||
{ # mutable-ok: prisma filter literal
|
||||
"OR": [ # mutable-ok: prisma filter literal
|
||||
{"team_id": None}, # mutable-ok: prisma filter literal
|
||||
{"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}}, # mutable-ok: prisma filter literal
|
||||
]
|
||||
},
|
||||
{ # mutable-ok: prisma filter literal
|
||||
"OR": [ # mutable-ok: prisma filter literal
|
||||
{"allowed_routes": {"is_empty": True}}, # mutable-ok: prisma filter literal
|
||||
{"allowed_routes": {"has": "llm_api_routes"}}, # mutable-ok: prisma filter literal
|
||||
]
|
||||
},
|
||||
],
|
||||
},
|
||||
order={"created_at": "desc"}, # mutable-ok: prisma order literal
|
||||
|
|
|
|||
|
|
@ -2199,6 +2199,43 @@ async def test_auto_register_map_existing_key_reuses_users_latest_key():
|
|||
assert result.org_id == "key-own-org"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_map_existing_key_skips_keys_that_cannot_call_llm_routes():
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
|
||||
return_value=SimpleNamespace(token="existing-hash")
|
||||
)
|
||||
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(
|
||||
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:
|
||||
await _auto_register_jwt_mapping(**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler))
|
||||
|
||||
where = prisma_client.db.litellm_verificationtoken.find_first.await_args.kwargs["where"]
|
||||
team_or: Final = next(entry["OR"] for entry in where["AND"] if any("team_id" in e for e in entry["OR"]))
|
||||
assert {"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}} in team_or, f"UI session keys must be excluded: {where}"
|
||||
assert {"team_id": None} in team_or, f"keys without a team must stay eligible: {where}"
|
||||
routes_or: Final = next(entry["OR"] for entry in where["AND"] if any("allowed_routes" in e for e in entry["OR"]))
|
||||
assert routes_or == [
|
||||
{"allowed_routes": {"is_empty": True}},
|
||||
{"allowed_routes": {"has": "llm_api_routes"}},
|
||||
], f"only unrestricted or llm_api_routes keys may be reused: {where}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_map_existing_key_mints_when_user_has_no_key():
|
||||
"""The flag must not leave a keyless user unmapped: with no existing key it
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue