mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
cache adjustment for deleted teams and keys
This commit is contained in:
parent
b06d7166e5
commit
c311033f48
2 changed files with 282 additions and 29 deletions
|
|
@ -11,7 +11,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_access_object,
|
||||
_cache_key_object,
|
||||
_cache_team_object,
|
||||
_delete_cache_access_object,
|
||||
_get_team_object_from_cache,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
|
@ -237,6 +240,10 @@ async def delete_access_group(
|
|||
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
# Track affected team IDs and key tokens for cache invalidation
|
||||
affected_team_ids: list = []
|
||||
affected_key_tokens: list = []
|
||||
|
||||
async with prisma_client.db.tx() as tx:
|
||||
existing = await tx.litellm_accessgrouptable.find_unique(
|
||||
where={"access_group_id": access_group_id}
|
||||
|
|
@ -252,6 +259,7 @@ async def delete_access_group(
|
|||
where={"access_group_ids": {"hasSome": [access_group_id]}}
|
||||
)
|
||||
for team in teams_with_group:
|
||||
affected_team_ids.append(team.team_id)
|
||||
updated_ids = [tid for tid in (team.access_group_ids or []) if tid != access_group_id]
|
||||
await tx.litellm_teamtable.update(
|
||||
where={"team_id": team.team_id},
|
||||
|
|
@ -262,6 +270,7 @@ async def delete_access_group(
|
|||
where={"access_group_ids": {"hasSome": [access_group_id]}}
|
||||
)
|
||||
for key in keys_with_group:
|
||||
affected_key_tokens.append(key.token)
|
||||
updated_ids = [kid for kid in (key.access_group_ids or []) if kid != access_group_id]
|
||||
await tx.litellm_verificationtoken.update(
|
||||
where={"token": key.token},
|
||||
|
|
@ -275,6 +284,44 @@ async def delete_access_group(
|
|||
# Invalidate the deleted access group from cache
|
||||
await _invalidate_cache_access_group(access_group_id)
|
||||
|
||||
# Patch cached team and key objects to remove the deleted access_group_id
|
||||
# instead of fully invalidating them (keeps cache warm, avoids DB re-fetch)
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
|
||||
|
||||
for team_id in affected_team_ids:
|
||||
cached_team = await _get_team_object_from_cache(
|
||||
key="team_id:{}".format(team_id),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
if cached_team is not None and cached_team.access_group_ids:
|
||||
cached_team.access_group_ids = [
|
||||
ag_id for ag_id in cached_team.access_group_ids if ag_id != access_group_id
|
||||
]
|
||||
await _cache_team_object(
|
||||
team_id=team_id,
|
||||
team_table=cached_team,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
for token in affected_key_tokens:
|
||||
cached_key = await user_api_key_cache.async_get_cache(key=token)
|
||||
if cached_key is not None:
|
||||
if isinstance(cached_key, dict):
|
||||
cached_key = UserAPIKeyAuth(**cached_key)
|
||||
if isinstance(cached_key, UserAPIKeyAuth) and cached_key.access_group_ids:
|
||||
cached_key.access_group_ids = [
|
||||
ag_id for ag_id in cached_key.access_group_ids if ag_id != access_group_id
|
||||
]
|
||||
await _cache_key_object(
|
||||
hashed_token=token,
|
||||
user_api_key_obj=cached_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -127,6 +127,7 @@ def client_and_mocks(monkeypatch):
|
|||
# Mock user_api_key_cache and proxy_logging_obj for cache operations (create/update/delete)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_set_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.delete_cache = MagicMock(return_value=None)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", mock_cache)
|
||||
|
||||
|
|
@ -136,6 +137,12 @@ def client_and_mocks(monkeypatch):
|
|||
mock_proxy_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_proxy_logging.internal_usage_cache.dual_cache.async_set_cache = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
monkeypatch.setattr(ps, "proxy_logging_obj", mock_proxy_logging)
|
||||
|
||||
admin_user = UserAPIKeyAuth(
|
||||
|
|
@ -146,7 +153,7 @@ def client_and_mocks(monkeypatch):
|
|||
|
||||
client = TestClient(app)
|
||||
|
||||
yield client, mock_prisma, mock_access_group_table
|
||||
yield client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging
|
||||
|
||||
app.dependency_overrides.clear()
|
||||
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
|
||||
|
|
@ -177,7 +184,7 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
|
|||
)
|
||||
def test_create_access_group_success(client_and_mocks, base_path, payload):
|
||||
"""Create access group with various payloads returns 201."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
resp = client.post(base_path, json=payload)
|
||||
assert resp.status_code == 201
|
||||
|
|
@ -189,7 +196,7 @@ def test_create_access_group_success(client_and_mocks, base_path, payload):
|
|||
|
||||
def test_create_access_group_duplicate_name_conflict(client_and_mocks):
|
||||
"""Create with duplicate name returns 409."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_name="existing-group")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -209,7 +216,7 @@ def test_create_access_group_duplicate_name_conflict(client_and_mocks):
|
|||
)
|
||||
def test_create_access_group_race_condition_returns_409(client_and_mocks, error_message):
|
||||
"""Create race condition: Prisma unique constraint surfaces as 409, not 500."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
mock_table.find_unique = AsyncMock(return_value=None)
|
||||
mock_table.create = AsyncMock(side_effect=Exception(error_message))
|
||||
|
|
@ -222,7 +229,7 @@ def test_create_access_group_race_condition_returns_409(client_and_mocks, error_
|
|||
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_create_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
||||
"""Non-admin users cannot create access groups."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="regular_user",
|
||||
|
|
@ -236,7 +243,7 @@ def test_create_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
|||
|
||||
def test_create_access_group_validation_missing_name(client_and_mocks):
|
||||
"""Create with missing access_group_name returns 422."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
resp = client.post("/v1/access_group", json={})
|
||||
assert resp.status_code == 422
|
||||
|
|
@ -244,7 +251,7 @@ def test_create_access_group_validation_missing_name(client_and_mocks):
|
|||
|
||||
def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks):
|
||||
"""Create with non-unique-constraint Prisma error returns 500."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
mock_table.find_unique = AsyncMock(return_value=None)
|
||||
mock_table.create = AsyncMock(side_effect=Exception("Some other database error"))
|
||||
|
|
@ -263,7 +270,7 @@ def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks
|
|||
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
|
||||
def test_list_access_groups_success_empty(client_and_mocks, base_path):
|
||||
"""List access groups returns empty list when none exist."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
resp = client.get(base_path)
|
||||
assert resp.status_code == 200
|
||||
|
|
@ -274,7 +281,7 @@ def test_list_access_groups_success_empty(client_and_mocks, base_path):
|
|||
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
|
||||
def test_list_access_groups_success_with_items(client_and_mocks, base_path):
|
||||
"""List access groups returns items when they exist."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
records = [
|
||||
_make_access_group_record(access_group_id="ag-1", access_group_name="group-1"),
|
||||
|
|
@ -293,7 +300,7 @@ def test_list_access_groups_success_with_items(client_and_mocks, base_path):
|
|||
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
|
||||
def test_list_access_groups_ordered_by_created_at_desc(client_and_mocks, base_path):
|
||||
"""List access groups calls find_many with created_at desc order."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
older = datetime(2025, 1, 1, 12, 0, 0)
|
||||
newer = datetime(2025, 1, 2, 12, 0, 0)
|
||||
|
|
@ -324,7 +331,7 @@ def test_list_access_groups_ordered_by_created_at_desc(client_and_mocks, base_pa
|
|||
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_list_access_groups_forbidden_non_admin(client_and_mocks, user_role):
|
||||
"""Non-admin users cannot list access groups."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="regular_user",
|
||||
|
|
@ -345,7 +352,7 @@ def test_list_access_groups_forbidden_non_admin(client_and_mocks, user_role):
|
|||
@pytest.mark.parametrize("access_group_id", ["ag-123", "ag-other-id"])
|
||||
def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
|
||||
"""Get access group by id returns record when found."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
record = _make_access_group_record(access_group_id=access_group_id)
|
||||
mock_table.find_unique = AsyncMock(return_value=record)
|
||||
|
|
@ -357,7 +364,7 @@ def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
|
|||
|
||||
def test_get_access_group_not_found(client_and_mocks):
|
||||
"""Get access group returns 404 when not found."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
mock_table.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -369,7 +376,7 @@ def test_get_access_group_not_found(client_and_mocks):
|
|||
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_get_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
||||
"""Non-admin users cannot get access group."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="regular_user",
|
||||
|
|
@ -397,7 +404,7 @@ def test_get_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
|||
)
|
||||
def test_update_access_group_success(client_and_mocks, base_path, update_payload):
|
||||
"""Update access group with various payloads returns 200."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -409,7 +416,7 @@ def test_update_access_group_success(client_and_mocks, base_path, update_payload
|
|||
|
||||
def test_update_access_group_not_found(client_and_mocks):
|
||||
"""Update access group returns 404 when not found."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
mock_table.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -425,7 +432,7 @@ def test_update_access_group_not_found(client_and_mocks):
|
|||
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_update_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
||||
"""Non-admin users cannot update access groups."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="regular_user",
|
||||
|
|
@ -439,7 +446,7 @@ def test_update_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
|||
|
||||
def test_update_access_group_empty_body(client_and_mocks):
|
||||
"""Update with empty body succeeds; only updated_by is set."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="unchanged")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -455,7 +462,7 @@ def test_update_access_group_empty_body(client_and_mocks):
|
|||
|
||||
def test_update_access_group_name_success(client_and_mocks):
|
||||
"""Update access_group_name succeeds when new name is unique."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -469,7 +476,7 @@ def test_update_access_group_name_success(client_and_mocks):
|
|||
|
||||
def test_update_access_group_name_duplicate_conflict(client_and_mocks):
|
||||
"""Update access_group_name to existing name returns 409 (unique constraint)."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -493,7 +500,7 @@ def test_update_access_group_name_duplicate_conflict(client_and_mocks):
|
|||
)
|
||||
def test_update_access_group_name_unique_constraint_returns_409(client_and_mocks, error_message):
|
||||
"""Update access_group_name: Prisma unique constraint surfaces as 409."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -513,7 +520,7 @@ def test_update_access_group_name_unique_constraint_returns_409(client_and_mocks
|
|||
@pytest.mark.parametrize("access_group_id", ["ag-123", "ag-delete-me"])
|
||||
def test_delete_access_group_success(client_and_mocks, base_path, access_group_id):
|
||||
"""Delete access group returns 204 when found."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id=access_group_id)
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -525,7 +532,7 @@ def test_delete_access_group_success(client_and_mocks, base_path, access_group_i
|
|||
|
||||
def test_delete_access_group_not_found(client_and_mocks):
|
||||
"""Delete access group returns 404 when not found."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
mock_table.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -538,7 +545,7 @@ def test_delete_access_group_not_found(client_and_mocks):
|
|||
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_delete_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
||||
"""Non-admin users cannot delete access groups."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="regular_user",
|
||||
|
|
@ -552,7 +559,7 @@ def test_delete_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
|||
|
||||
def test_delete_access_group_cleans_up_teams_and_keys(client_and_mocks):
|
||||
"""Delete removes access_group_id from teams and keys before deleting the group."""
|
||||
client, mock_prisma, mock_access_group_table = client_and_mocks
|
||||
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
|
||||
mock_team_table = mock_prisma.db.litellm_teamtable
|
||||
mock_key_table = mock_prisma.db.litellm_verificationtoken
|
||||
|
||||
|
|
@ -585,9 +592,208 @@ def test_delete_access_group_cleans_up_teams_and_keys(client_and_mocks):
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"team_cache_group_ids,key_cache_group_ids,expected_team_ids_after,expected_key_ids_after",
|
||||
[
|
||||
# Team and key both cached with the deleted group
|
||||
(
|
||||
["ag-to-delete", "ag-keep"],
|
||||
["ag-to-delete", "ag-stay"],
|
||||
["ag-keep"],
|
||||
["ag-stay"],
|
||||
),
|
||||
# Only team cached; key not in cache
|
||||
(
|
||||
["ag-to-delete"],
|
||||
None,
|
||||
[],
|
||||
None,
|
||||
),
|
||||
# Only key cached; team not in cache
|
||||
(
|
||||
None,
|
||||
["ag-to-delete"],
|
||||
None,
|
||||
[],
|
||||
),
|
||||
# Neither cached — nothing to patch
|
||||
(
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
),
|
||||
# Cached team has only the deleted group
|
||||
(
|
||||
["ag-to-delete"],
|
||||
["ag-to-delete"],
|
||||
[],
|
||||
[],
|
||||
),
|
||||
# Cached objects have multiple groups, only the deleted one is removed
|
||||
(
|
||||
["ag-alpha", "ag-to-delete", "ag-beta"],
|
||||
["ag-to-delete", "ag-gamma"],
|
||||
["ag-alpha", "ag-beta"],
|
||||
["ag-gamma"],
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"both_cached",
|
||||
"only_team_cached",
|
||||
"only_key_cached",
|
||||
"neither_cached",
|
||||
"single_group_removed",
|
||||
"multi_group_partial_removal",
|
||||
],
|
||||
)
|
||||
def test_delete_access_group_patches_cached_team_and_key(
|
||||
client_and_mocks,
|
||||
team_cache_group_ids,
|
||||
key_cache_group_ids,
|
||||
expected_team_ids_after,
|
||||
expected_key_ids_after,
|
||||
):
|
||||
"""Delete patches cached team/key objects to remove the deleted access_group_id."""
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
|
||||
mock_team_table = mock_prisma.db.litellm_teamtable
|
||||
mock_key_table = mock_prisma.db.litellm_verificationtoken
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-to-delete")
|
||||
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
# Set up a team and key in the DB that reference the group
|
||||
team_with_group = MagicMock()
|
||||
team_with_group.team_id = "team-1"
|
||||
team_with_group.access_group_ids = ["ag-to-delete", "ag-keep"]
|
||||
mock_team_table.find_many = AsyncMock(return_value=[team_with_group])
|
||||
|
||||
key_with_group = MagicMock()
|
||||
key_with_group.token = "hashed-key-1"
|
||||
key_with_group.access_group_ids = ["ag-to-delete"]
|
||||
mock_key_table.find_many = AsyncMock(return_value=[key_with_group])
|
||||
|
||||
# Build cached team object (returned from proxy_logging dual cache)
|
||||
if team_cache_group_ids is not None:
|
||||
cached_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team-1",
|
||||
access_group_ids=list(team_cache_group_ids),
|
||||
)
|
||||
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
|
||||
return_value=cached_team
|
||||
)
|
||||
else:
|
||||
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
# Build cached key object (returned from user_api_key_cache)
|
||||
if key_cache_group_ids is not None:
|
||||
cached_key = UserAPIKeyAuth(
|
||||
token="hashed-key-1",
|
||||
access_group_ids=list(key_cache_group_ids),
|
||||
)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=cached_key)
|
||||
else:
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
||||
resp = client.delete("/v1/access_group/ag-to-delete")
|
||||
assert resp.status_code == 204
|
||||
|
||||
# Verify DB cleanup always happens
|
||||
mock_team_table.update.assert_awaited_once()
|
||||
mock_key_table.update.assert_awaited_once()
|
||||
|
||||
# Verify cache patching
|
||||
if expected_team_ids_after is not None:
|
||||
# _cache_team_object writes via _cache_management_object -> async_set_cache
|
||||
team_set_calls = [
|
||||
c for c in mock_cache.async_set_cache.call_args_list
|
||||
if c.kwargs.get("key", "") == "team_id:team-1"
|
||||
or (len(c.args) >= 1 and c.args[0] == "team_id:team-1")
|
||||
]
|
||||
assert len(team_set_calls) >= 1, "Expected team cache to be patched"
|
||||
# The cached team object should have the updated access_group_ids
|
||||
written_team = team_set_calls[0].kwargs.get("value") or team_set_calls[0].args[1]
|
||||
if isinstance(written_team, LiteLLM_TeamTableCachedObj):
|
||||
assert written_team.access_group_ids == expected_team_ids_after
|
||||
else:
|
||||
# No team in cache — async_set_cache should not be called for team_id key
|
||||
team_set_calls = [
|
||||
c for c in mock_cache.async_set_cache.call_args_list
|
||||
if c.kwargs.get("key", "") == "team_id:team-1"
|
||||
or (len(c.args) >= 1 and c.args[0] == "team_id:team-1")
|
||||
]
|
||||
assert len(team_set_calls) == 0, "Should not patch team cache when not cached"
|
||||
|
||||
if expected_key_ids_after is not None:
|
||||
key_set_calls = [
|
||||
c for c in mock_cache.async_set_cache.call_args_list
|
||||
if c.kwargs.get("key", "") == "hashed-key-1"
|
||||
or (len(c.args) >= 1 and c.args[0] == "hashed-key-1")
|
||||
]
|
||||
assert len(key_set_calls) >= 1, "Expected key cache to be patched"
|
||||
written_key = key_set_calls[0].kwargs.get("value") or key_set_calls[0].args[1]
|
||||
if isinstance(written_key, UserAPIKeyAuth):
|
||||
assert written_key.access_group_ids == expected_key_ids_after
|
||||
else:
|
||||
key_set_calls = [
|
||||
c for c in mock_cache.async_set_cache.call_args_list
|
||||
if c.kwargs.get("key", "") == "hashed-key-1"
|
||||
or (len(c.args) >= 1 and c.args[0] == "hashed-key-1")
|
||||
]
|
||||
assert len(key_set_calls) == 0, "Should not patch key cache when not cached"
|
||||
|
||||
|
||||
def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks):
|
||||
"""Delete correctly patches a key cached as a raw dict (not UserAPIKeyAuth)."""
|
||||
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
|
||||
mock_team_table = mock_prisma.db.litellm_teamtable
|
||||
mock_key_table = mock_prisma.db.litellm_verificationtoken
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-to-delete")
|
||||
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
mock_team_table.find_many = AsyncMock(return_value=[])
|
||||
|
||||
key_with_group = MagicMock()
|
||||
key_with_group.token = "hashed-key-dict"
|
||||
key_with_group.access_group_ids = ["ag-to-delete", "ag-other"]
|
||||
mock_key_table.find_many = AsyncMock(return_value=[key_with_group])
|
||||
|
||||
# No team in cache
|
||||
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
# Key cached as a plain dict (as can happen with Redis serialization)
|
||||
mock_cache.async_get_cache = AsyncMock(
|
||||
return_value={
|
||||
"token": "hashed-key-dict",
|
||||
"access_group_ids": ["ag-to-delete", "ag-other"],
|
||||
}
|
||||
)
|
||||
|
||||
resp = client.delete("/v1/access_group/ag-to-delete")
|
||||
assert resp.status_code == 204
|
||||
|
||||
# The key should have been re-cached with the deleted group removed
|
||||
key_set_calls = [
|
||||
c for c in mock_cache.async_set_cache.call_args_list
|
||||
if c.kwargs.get("key", "") == "hashed-key-dict"
|
||||
or (len(c.args) >= 1 and c.args[0] == "hashed-key-dict")
|
||||
]
|
||||
assert len(key_set_calls) >= 1, "Expected key cache to be patched"
|
||||
written_key = key_set_calls[0].kwargs.get("value") or key_set_calls[0].args[1]
|
||||
if isinstance(written_key, UserAPIKeyAuth):
|
||||
assert written_key.access_group_ids == ["ag-other"]
|
||||
|
||||
|
||||
def test_delete_access_group_503_on_db_connection_error(client_and_mocks):
|
||||
"""Delete returns 503 when DB connection error occurs during transaction."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-to-delete")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -600,7 +806,7 @@ def test_delete_access_group_503_on_db_connection_error(client_and_mocks):
|
|||
|
||||
def test_delete_access_group_404_on_p2025_or_record_not_found(client_and_mocks):
|
||||
"""Delete returns 404 when Prisma raises P2025 or record-not-found error."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-to-delete")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -613,7 +819,7 @@ def test_delete_access_group_404_on_p2025_or_record_not_found(client_and_mocks):
|
|||
|
||||
def test_delete_access_group_500_on_generic_exception(client_and_mocks):
|
||||
"""Delete returns 500 when generic exception occurs during transaction."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
client, _, mock_table, *_ = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-to-delete")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
|
@ -647,7 +853,7 @@ def test_delete_access_group_500_on_generic_exception(client_and_mocks):
|
|||
)
|
||||
def test_access_group_endpoints_db_not_connected(client_and_mocks, monkeypatch, method, url, factory):
|
||||
"""All endpoints return 500 when DB is not connected."""
|
||||
client, _, _ = client_and_mocks
|
||||
client, *_ = client_and_mocks
|
||||
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue