mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): evict jwt_key_mapping cache when user, team, or org deletion removes mapped keys
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
10616d7407
commit
884a467a7f
8 changed files with 462 additions and 11 deletions
|
|
@ -3637,6 +3637,22 @@ async def get_jwt_key_mapping_cache_keys_for_token(
|
|||
return tuple(jwt_key_mapping_cache_key(m.jwt_claim_name, m.jwt_claim_value, m.jwt_issuer) for m in mappings)
|
||||
|
||||
|
||||
class _TokenInFilter(TypedDict):
|
||||
token: ReadOnly[Mapping[str, Sequence[str]]]
|
||||
|
||||
|
||||
async def get_jwt_key_mapping_cache_keys_for_tokens(
|
||||
hashed_tokens: Sequence[str],
|
||||
prisma_client: PrismaClient,
|
||||
) -> tuple[str, ...]:
|
||||
"""Cache keys of every JWT claim mapped to any of the given virtual keys."""
|
||||
if not hashed_tokens:
|
||||
return ()
|
||||
token_filter: Final[_TokenInFilter] = {"token": {"in": tuple(hashed_tokens)}}
|
||||
mappings: Final = await _jwt_key_mapping_table(JWTKeyMappingRepository(prisma_client)).find_many(where=token_filter)
|
||||
return tuple(jwt_key_mapping_cache_key(m.jwt_claim_name, m.jwt_claim_value, m.jwt_issuer) for m in mappings)
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_jwt_key_mapping_object(
|
||||
jwt_claim_name: str,
|
||||
|
|
|
|||
|
|
@ -23,12 +23,18 @@ from typing import Any, Final, Literal, Protocol, cast, overload
|
|||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
delete_cache_key_objects,
|
||||
get_jwt_key_mapping_cache_keys_for_tokens,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.auth.password_policy import validate_password_policy
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
|
|
@ -126,6 +132,10 @@ def _verification_token_table(
|
|||
return token_table
|
||||
|
||||
|
||||
class _UserIdInFilter(TypedDict):
|
||||
user_id: ReadOnly[Mapping[str, Sequence[str]]]
|
||||
|
||||
|
||||
def _organization_membership_table(
|
||||
prisma_client: "PrismaClient | None",
|
||||
) -> "TableActions[prisma_models.LiteLLM_OrganizationMembership]":
|
||||
|
|
@ -2345,6 +2355,8 @@ async def delete_user(
|
|||
create_audit_log_for_update,
|
||||
litellm_proxy_admin_name,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -2471,7 +2483,20 @@ async def delete_user(
|
|||
# End of Audit logging
|
||||
|
||||
## DELETE ASSOCIATED KEYS
|
||||
await _verification_token_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
|
||||
key_filter: Final[_UserIdInFilter] = {"user_id": {"in": data.user_ids}}
|
||||
keys_to_delete: Final = await _verification_token_table(prisma_client).find_many(where=key_filter)
|
||||
hashed_tokens_to_delete: Final = tuple(key.token for key in keys_to_delete)
|
||||
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
|
||||
hashed_tokens=hashed_tokens_to_delete,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
await _verification_token_table(prisma_client).delete_many(where=key_filter)
|
||||
await delete_cache_key_objects(
|
||||
hashed_tokens=hashed_tokens_to_delete,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
|
||||
|
||||
## DELETE ASSOCIATED INVITATION LINKS
|
||||
await _invitation_link_table(prisma_client).delete_many(
|
||||
|
|
|
|||
|
|
@ -26,13 +26,21 @@ from typing import (
|
|||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import can_user_call_model, get_user_object
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
can_user_call_model,
|
||||
delete_cache_key_objects,
|
||||
get_jwt_key_mapping_cache_keys_for_tokens,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.budget_management_endpoints import (
|
||||
new_budget,
|
||||
update_budget,
|
||||
|
|
@ -52,7 +60,7 @@ from litellm.proxy.management_helpers.utils import (
|
|||
get_new_internal_user_defaults,
|
||||
management_endpoint_wrapper,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
|
|
@ -79,6 +87,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from prisma.models import LiteLLM_OrganizationTable as PrismaOrganizationTable
|
||||
from prisma.models import LiteLLM_UserTable as PrismaUserTable
|
||||
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationToken
|
||||
|
||||
|
||||
async def _enterprise_license_required(
|
||||
|
|
@ -168,9 +177,15 @@ class _TeamTableClient(Protocol):
|
|||
|
||||
|
||||
class _VerificationTokenTableClient(Protocol):
|
||||
async def find_many(self, where: Mapping[str, object] | None = None) -> "Sequence[PrismaVerificationToken]": ...
|
||||
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _OrganizationIdFilter(TypedDict):
|
||||
organization_id: ReadOnly[str]
|
||||
|
||||
|
||||
class _ObjectPermissionTxClient(Protocol):
|
||||
async def upsert(
|
||||
self, where: Mapping[str, object], data: Mapping[str, object]
|
||||
|
|
@ -961,7 +976,7 @@ async def delete_organization(
|
|||
|
||||
- organization_ids: List[str] - The organization ids to delete.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -983,8 +998,12 @@ async def delete_organization(
|
|||
await _table(OrganizationMembershipRepository(prisma_client)).delete_many(
|
||||
where={"organization_id": organization_id}
|
||||
)
|
||||
# delete all keys in the organization
|
||||
await _table(VerificationTokenRepository(prisma_client)).delete_many(where={"organization_id": organization_id})
|
||||
await _delete_organization_keys(
|
||||
organization_id=organization_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
# delete the organization
|
||||
deleted_org = await _table(OrganizationRepository(prisma_client)).delete(
|
||||
where={"organization_id": organization_id},
|
||||
|
|
@ -1000,6 +1019,30 @@ async def delete_organization(
|
|||
return deleted_orgs
|
||||
|
||||
|
||||
async def _delete_organization_keys(
|
||||
organization_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
) -> None:
|
||||
"""Delete the organization's keys and drop every cache entry that still resolves to them,
|
||||
including the jwt_key_mapping entries the FK cascade removes from the table but not from the cache."""
|
||||
key_filter: Final[_OrganizationIdFilter] = {"organization_id": organization_id}
|
||||
keys_to_delete: Final = await _table(VerificationTokenRepository(prisma_client)).find_many(where=key_filter)
|
||||
hashed_tokens_to_delete: Final = tuple(key.token for key in keys_to_delete)
|
||||
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
|
||||
hashed_tokens=hashed_tokens_to_delete,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
await _table(VerificationTokenRepository(prisma_client)).delete_many(where=key_filter)
|
||||
await delete_cache_key_objects(
|
||||
hashed_tokens=hashed_tokens_to_delete,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/organization/list",
|
||||
tags=["organization management"],
|
||||
|
|
|
|||
|
|
@ -99,6 +99,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
can_org_access_model,
|
||||
delete_cache_key_objects,
|
||||
delete_cache_team_object,
|
||||
get_jwt_key_mapping_cache_keys_for_tokens,
|
||||
get_org_object,
|
||||
get_team_membership,
|
||||
get_team_object,
|
||||
|
|
@ -110,6 +111,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
enforce_output_token_estimates_are_admin_only,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
|
@ -3524,7 +3526,6 @@ async def team_member_delete(
|
|||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -3626,6 +3627,10 @@ async def team_member_delete(
|
|||
"team_id": data.team_id,
|
||||
}
|
||||
)
|
||||
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
|
||||
hashed_tokens=tuple(key.token for key in keys_to_delete),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if removed_team_members:
|
||||
await _team_tx_db(tx).update(
|
||||
|
|
@ -3674,6 +3679,7 @@ async def team_member_delete(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
|
||||
await evict_and_broadcast(cache_keys=tuple(sorted(user_ids_to_delete)), user_api_key_cache=user_api_key_cache)
|
||||
for user_id in sorted(user_ids_to_delete):
|
||||
await invalidate_team_member_spend_state(
|
||||
|
|
@ -4264,6 +4270,10 @@ async def delete_team(
|
|||
)
|
||||
|
||||
keys_to_delete: Final = await _tokens_db(prisma_client).find_many(where={"team_id": {"in": data.team_ids}})
|
||||
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
|
||||
hashed_tokens=tuple(key.token for key in keys_to_delete),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if keys_to_delete:
|
||||
await _persist_deleted_verification_tokens(
|
||||
|
|
@ -4280,6 +4290,7 @@ async def delete_team(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
|
||||
|
||||
## DELETE ASSOCIATED BYOK MODELS
|
||||
# Runs before the team rows are deleted so a mid-flight failure never leaves
|
||||
|
|
|
|||
|
|
@ -2627,6 +2627,9 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
|
|||
)
|
||||
|
||||
# Mock all delete_many calls
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = mocker.AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.delete_many = mocker.AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
|
|
@ -2676,6 +2679,105 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
|
|||
assert condition[field] == {"in": ["admin-creator"]}
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker):
|
||||
"""/user/delete bulk-deletes the user's keys without going through /key/delete, so the
|
||||
jwt_key_mapping cache entries pointing at those keys must be evicted here too. A surviving
|
||||
entry keeps resolving the deleted token hash until the mapping cache TTL expires: the deleted
|
||||
identity is either still served through the stale key cache or 401s on every JWT call, and it is
|
||||
never re-registered (LIT-5387).
|
||||
|
||||
The FK cascade drops the mapping rows with the key rows, so the cache keys have to be read
|
||||
before the delete: reading them afterwards finds nothing to evict.
|
||||
"""
|
||||
from litellm.proxy._types import DeleteUserRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
|
||||
|
||||
global_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", None)
|
||||
issuer_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", "https://issuer.example")
|
||||
unrelated_cache_key: Final = jwt_key_mapping_cache_key("sub", "other-user", None)
|
||||
jwt_table: Final = _CascadingJWTMappingTable(
|
||||
[
|
||||
_JWTMappingRow("hashed-jwt-key", "sub", "jwt-user"),
|
||||
_JWTMappingRow("hashed-issuer-key", "sub", "jwt-user", "https://issuer.example"),
|
||||
_JWTMappingRow("hashed-unrelated-key", "sub", "other-user"),
|
||||
]
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
for cache_key, hashed_token in (
|
||||
(global_cache_key, "hashed-jwt-key"),
|
||||
(issuer_cache_key, "hashed-issuer-key"),
|
||||
(unrelated_cache_key, "hashed-unrelated-key"),
|
||||
):
|
||||
cache.set_cache(key=cache_key, value=hashed_token)
|
||||
cache.set_cache(key=hashed_token, value=UserAPIKeyAuth(token=hashed_token))
|
||||
|
||||
user_row: Final = mocker.MagicMock()
|
||||
user_row.user_id = "jwt-user"
|
||||
user_row.user_email = "jwt-user@example.com"
|
||||
user_row.teams = []
|
||||
user_row.model_dump_json.return_value = "{}"
|
||||
user_row.model_dump.return_value = {"user_id": "jwt-user", "user_email": "jwt-user@example.com", "teams": []}
|
||||
|
||||
mock_prisma_client: Final = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=user_row)
|
||||
mock_prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = mocker.AsyncMock(
|
||||
return_value=[SimpleNamespace(token="hashed-jwt-key"), SimpleNamespace(token="hashed-issuer-key")]
|
||||
)
|
||||
|
||||
async def cascading_delete_many(where):
|
||||
jwt_table.cascade(("hashed-jwt-key", "hashed-issuer-key"))
|
||||
return 2
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.delete_many = mocker.AsyncMock(side_effect=cascading_delete_many)
|
||||
mock_prisma_client.db.litellm_invitationlink.delete_many = mocker.AsyncMock(return_value=0)
|
||||
mock_prisma_client.db.litellm_organizationmembership.delete_many = mocker.AsyncMock(return_value=0)
|
||||
mock_prisma_client.db.litellm_teammembership.delete_many = mocker.AsyncMock(return_value=0)
|
||||
mock_prisma_client.db.litellm_usertable.delete_many = mocker.AsyncMock(return_value=1)
|
||||
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # test-quality-ok: substitute the database dependency
|
||||
mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: exercise a real isolated cache
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", None) # test-quality-ok: delete_user reads it off proxy_server at call time
|
||||
|
||||
await delete_user(
|
||||
data=DeleteUserRequest(user_ids=["jwt-user"]),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="proxy-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert cache.get_cache(key=global_cache_key) is None
|
||||
assert cache.get_cache(key=issuer_cache_key) is None
|
||||
assert cache.get_cache(key="hashed-jwt-key") is None
|
||||
assert cache.get_cache(key="hashed-issuer-key") is None
|
||||
assert cache.get_cache(key=unrelated_cache_key) == "hashed-unrelated-key"
|
||||
assert cache.get_cache(key="hashed-unrelated-key") is not None
|
||||
assert [row.token for row in jwt_table.rows] == ["hashed-unrelated-key"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker):
|
||||
"""Regression: an org admin of org-A must not be able to delete a user
|
||||
|
|
|
|||
|
|
@ -5244,7 +5244,10 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat
|
|||
virtual_key_mapping_cache_ttl expires, instead of auto-registering again.
|
||||
"""
|
||||
jwt_table = _CascadingJWTMappingTable(
|
||||
[_JWTMappingRow("hashed-token-1", "email", "user@example.com")]
|
||||
[
|
||||
_JWTMappingRow("hashed-token-1", "email", "user@example.com"),
|
||||
_JWTMappingRow("hashed-token-1", "email", "user@example.com", "https://issuer.example"),
|
||||
]
|
||||
)
|
||||
|
||||
key1 = LiteLLM_VerificationToken(
|
||||
|
|
@ -5302,7 +5305,10 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat
|
|||
),
|
||||
)
|
||||
|
||||
assert recording_evict.cache_keys == (jwt_key_mapping_cache_key("email", "user@example.com", None),)
|
||||
assert recording_evict.cache_keys == (
|
||||
jwt_key_mapping_cache_key("email", "user@example.com", None),
|
||||
jwt_key_mapping_cache_key("email", "user@example.com", "https://issuer.example"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
from litellm._uuid import uuid
|
||||
from types import MappingProxyType
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final, Mapping, Optional, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -1438,3 +1438,82 @@ def test_organization_routes_reach_their_handler_with_enterprise_license(monkeyp
|
|||
assert any(
|
||||
message in response.text for message in (CommonProxyErrors.db_not_connected_error.value, "No db connected")
|
||||
)
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monkeypatch):
|
||||
"""/organization/delete bulk-deletes the org's keys without going through /key/delete, so the
|
||||
key objects and the jwt_key_mapping entries (issuer-scoped ones included) pointing at them
|
||||
must be evicted here, or a deleted key keeps authenticating and a JWT identity keeps resolving
|
||||
a token hash that no longer exists until the TTLs expire. The FK cascade drops the mapping
|
||||
rows with the key rows, so the cache keys have to be read before the delete (LIT-5387)."""
|
||||
from litellm.proxy._types import DeleteOrganizationRequest, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import delete_organization
|
||||
|
||||
doomed_cache_keys: Final = (
|
||||
"hashed-org-key",
|
||||
jwt_key_mapping_cache_key("sub", "svc-account", None),
|
||||
jwt_key_mapping_cache_key("sub", "svc-account", "https://issuer.example"),
|
||||
)
|
||||
kept_cache_keys: Final = ("hashed-other-key", jwt_key_mapping_cache_key("sub", "other-account", None))
|
||||
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "other-account")
|
||||
jwt_table: Final = _CascadingJWTMappingTable(
|
||||
[
|
||||
_JWTMappingRow("hashed-org-key", "sub", "svc-account"),
|
||||
_JWTMappingRow("hashed-org-key", "sub", "svc-account", "https://issuer.example"),
|
||||
kept_row,
|
||||
]
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
for cache_key in (*doomed_cache_keys, *kept_cache_keys):
|
||||
cache.set_cache(key=cache_key, value={"retained": True})
|
||||
|
||||
prisma_client: Final = AsyncMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(token="hashed-org-key")]
|
||||
)
|
||||
|
||||
async def cascading_delete_many(where):
|
||||
jwt_table.cascade(("hashed-org-key",))
|
||||
return 1
|
||||
|
||||
prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many)
|
||||
prisma_client.db.litellm_jwtkeymapping = jwt_table
|
||||
prisma_client.db.litellm_organizationtable.delete = AsyncMock(return_value=MagicMock())
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True, raising=False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None)
|
||||
|
||||
await delete_organization(
|
||||
data=DeleteOrganizationRequest(organization_ids=["org-doomed"]),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert all(cache.get_cache(key=cache_key) == {"retained": True} for cache_key in kept_cache_keys)
|
||||
assert jwt_table.rows == [kept_row]
|
||||
|
|
|
|||
|
|
@ -9180,6 +9180,175 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch):
|
|||
assert cache.get_cache(key="unrelated-key") == {"retained": True}
|
||||
|
||||
|
||||
class _JWTMappingRow:
|
||||
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
|
||||
self.token = token
|
||||
self.jwt_claim_name = jwt_claim_name
|
||||
self.jwt_claim_value = jwt_claim_value
|
||||
self.jwt_issuer = jwt_issuer
|
||||
|
||||
|
||||
class _CascadingJWTMappingTable:
|
||||
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.rows = rows
|
||||
|
||||
async def find_many(self, where, **kwargs):
|
||||
return [row for row in self.rows if row.token in where["token"]["in"]]
|
||||
|
||||
def cascade(self, deleted_tokens):
|
||||
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
|
||||
|
||||
|
||||
def _seed_jwt_mapping_cache(cache, mapping_rows):
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
|
||||
cache_keys = tuple(
|
||||
jwt_key_mapping_cache_key(row.jwt_claim_name, row.jwt_claim_value, row.jwt_issuer) for row in mapping_rows
|
||||
)
|
||||
for cache_key, row in zip(cache_keys, mapping_rows):
|
||||
cache.set_cache(key=cache_key, value=row.token)
|
||||
return cache_keys
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(monkeypatch):
|
||||
"""The member's team keys are deleted in bulk here, not through /key/delete, so the
|
||||
jwt_key_mapping cache entries pointing at them must be evicted here too, or every JWT call
|
||||
from that identity resolves the deleted token hash and 401s until the mapping TTL expires.
|
||||
The FK cascade drops the mapping rows with the key rows, so the cache keys have to be read
|
||||
before the delete (LIT-5387)."""
|
||||
from litellm.proxy._types import TeamMemberDeleteRequest
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import LiteLLM_VerificationToken
|
||||
|
||||
doomed_rows: Final = (
|
||||
_JWTMappingRow("hashed-token-1", "sub", "user-123"),
|
||||
_JWTMappingRow("hashed-token-1", "sub", "user-123", "https://issuer.example"),
|
||||
)
|
||||
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "user-999")
|
||||
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-1",
|
||||
team_alias="test-team",
|
||||
members_with_roles=[Member(user_id="user-123", role="admin")],
|
||||
metadata={},
|
||||
model_max_budget={},
|
||||
model_spend={},
|
||||
)
|
||||
key1 = LiteLLM_VerificationToken(token="hashed-token-1", user_id="user-123", team_id="team-1")
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[MagicMock(user_id="user-123", teams=["team-1"])]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[key1])
|
||||
|
||||
async def cascading_delete_many(where):
|
||||
jwt_table.cascade(("hashed-token-1",))
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many)
|
||||
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
|
||||
_wire_member_delete_tx(mock_prisma_client)
|
||||
|
||||
cache: Final = UserApiKeyCache()
|
||||
doomed_cache_keys: Final = _seed_jwt_mapping_cache(cache, doomed_rows)
|
||||
(kept_cache_key,) = _seed_jwt_mapping_cache(cache, (kept_row,))
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
monkeypatch.setattr("litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin", lambda **kwargs: True)
|
||||
|
||||
await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id="team-1", user_id="user-123"),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="admin-user", api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value
|
||||
),
|
||||
)
|
||||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert cache.get_cache(key=kept_cache_key) == "hashed-other-key"
|
||||
assert jwt_table.rows == [kept_row]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_team_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(
|
||||
monkeypatch,
|
||||
disable_audit_logging_for_mocked_team,
|
||||
):
|
||||
"""Same contract as /team/member_delete for the bulk key delete in /team/delete: the
|
||||
jwt_key_mapping cache entries of the team's keys, issuer-scoped ones included, are gone
|
||||
after the delete while entries pointing at other keys survive (LIT-5387)."""
|
||||
from litellm.proxy._types import DeleteTeamRequest, LiteLLM_VerificationToken
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
doomed_rows: Final = (
|
||||
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account"),
|
||||
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account", "https://issuer.example"),
|
||||
)
|
||||
kept_row: Final = _JWTMappingRow("hashed-unrelated-key", "sub", "svc-account", "https://other-issuer.example")
|
||||
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-doomed",
|
||||
team_alias="doomed-team",
|
||||
members_with_roles=[],
|
||||
metadata={},
|
||||
model_max_budget={},
|
||||
model_spend={},
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
|
||||
async def cascading_delete_data(team_id_list, table_name):
|
||||
jwt_table.cascade(("hashed-doomed-key",))
|
||||
return {"deleted_keys": 1}
|
||||
|
||||
mock_prisma_client.delete_data = AsyncMock(side_effect=cascading_delete_data)
|
||||
mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock()
|
||||
mock_prisma_client.db.litellm_deletedverificationtoken.create_many = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[LiteLLM_VerificationToken(token="hashed-doomed-key", team_id="team-doomed")]
|
||||
)
|
||||
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
|
||||
mock_prisma_client.db.execute_raw = AsyncMock()
|
||||
mock_prisma_client.db.litellm_teammembership.delete_many = AsyncMock()
|
||||
|
||||
mock_tx = AsyncMock()
|
||||
mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
mock_tx_cm = MagicMock()
|
||||
mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx)
|
||||
mock_tx_cm.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm)
|
||||
_wire_team_delete_tx(mock_prisma_client)
|
||||
|
||||
cache: Final = UserApiKeyCache()
|
||||
doomed_cache_keys: Final = _seed_jwt_mapping_cache(cache, doomed_rows)
|
||||
(kept_cache_key,) = _seed_jwt_mapping_cache(cache, (kept_row,))
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.create_audit_log_for_update", AsyncMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")
|
||||
|
||||
await delete_team(
|
||||
data=DeleteTeamRequest(team_ids=["team-doomed"]),
|
||||
http_request=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="admin-user", api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value
|
||||
),
|
||||
litellm_changed_by="admin-user",
|
||||
)
|
||||
|
||||
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
|
||||
assert cache.get_cache(key=kept_cache_key) == "hashed-unrelated-key"
|
||||
assert jwt_table.rows == [kept_row]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_negative_max_budget():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue