mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(team): invalidate stale membership cache after bulk member update
get_team_membership caches each member's LiteLLM_TeamMembership under
team_membership:{user_id}:{team_id} with no TTL, and the legacy spend-check
path in user_api_key_auth.py caches the same row under {team_id}_{user_id}.
bulk_update_team_members only refreshed the team-row cache after committing
budget/allowed_models writes, so a member could keep invoking a newly
disallowed model or spending against a stale limit until those entries
separately expired. Evict both cache keys for every member the transaction
actually wrote.
This commit is contained in:
parent
1d09fd3ecf
commit
7fff840574
2 changed files with 113 additions and 8 deletions
|
|
@ -180,6 +180,27 @@ async def _refresh_cached_team(
|
|||
)
|
||||
|
||||
|
||||
async def _invalidate_team_membership_caches(
|
||||
team_id: str,
|
||||
user_ids: List[str],
|
||||
user_api_key_cache: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Invalidate the per-member membership caches after a budget/allowed_models
|
||||
write to `litellm_teammembership` for `user_ids` in `team_id`.
|
||||
|
||||
`get_team_membership` (auth_checks.py) caches `LiteLLM_TeamMembership`
|
||||
under `team_membership:{user_id}:{team_id}` with no TTL, and the legacy
|
||||
spend-check path in `user_api_key_auth.py` caches the same row under
|
||||
`{team_id}_{user_id}`. Without invalidating both, a member keeps their
|
||||
prior `allowed_models`/budget until the cache entry is separately evicted,
|
||||
bypassing restrictions an admin just applied.
|
||||
"""
|
||||
for user_id in user_ids:
|
||||
await user_api_key_cache.async_delete_cache(key="team_membership:{}:{}".format(user_id, team_id))
|
||||
await user_api_key_cache.async_delete_cache(key="{}_{}".format(team_id, user_id))
|
||||
|
||||
|
||||
async def _verify_team_access(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -3175,6 +3196,13 @@ async def bulk_update_team_members(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if budget_target_user_ids:
|
||||
await _invalidate_team_membership_caches(
|
||||
team_id=team_id,
|
||||
user_ids=budget_target_user_ids,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
successful_updates = [
|
||||
TeamMemberUpdateResponse(
|
||||
team_id=team_id,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.proxy._types import (
|
|||
TeamMemberUpdateRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_members,
|
||||
team_member_update,
|
||||
|
|
@ -217,15 +218,16 @@ class _FakeBulkDb:
|
|||
|
||||
def _bulk_setup(monkeypatch, team_row, memberships):
|
||||
db = _FakeBulkDb(team_row, memberships)
|
||||
user_api_key_cache = UserApiKeyCache()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", types.SimpleNamespace(db=db))
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
|
||||
upsert_mock = AsyncMock()
|
||||
monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock)
|
||||
refresh_mock = AsyncMock()
|
||||
monkeypatch.setattr(team_endpoints, "_refresh_cached_team", refresh_mock)
|
||||
return db, upsert_mock, refresh_mock
|
||||
return db, upsert_mock, refresh_mock, user_api_key_cache
|
||||
|
||||
|
||||
def _bulk_request():
|
||||
|
|
@ -249,7 +251,7 @@ async def test_bulk_update_delegates_budget_upsert_per_valid_member(monkeypatch)
|
|||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
)
|
||||
db, upsert_mock, refresh_mock = _bulk_setup(
|
||||
db, upsert_mock, refresh_mock, _cache = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[
|
||||
|
|
@ -282,6 +284,81 @@ async def test_bulk_update_delegates_budget_upsert_per_valid_member(monkeypatch)
|
|||
assert [member.user_id for member in response.successful_updates] == ["user-1", "user-2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_invalidates_membership_cache_for_written_members_only(monkeypatch):
|
||||
"""A budget/allowed_models write must evict the per-member
|
||||
`team_membership:{user_id}:{team_id}` cache and the legacy
|
||||
`{team_id}_{user_id}` cache for every member actually written, so the next
|
||||
request re-reads the new limits instead of enforcing stale ones. Members
|
||||
the patch didn't touch (no budget_patch, or filtered out as non-members)
|
||||
must be left alone."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
Member(user_id="user-1", role="user"),
|
||||
Member(user_id="user-2", role="user"),
|
||||
],
|
||||
)
|
||||
_db, _upsert_mock, _refresh_mock, cache = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[
|
||||
LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1"),
|
||||
LiteLLM_TeamMembership(user_id="user-2", team_id="team-1234", budget_id="bud-2"),
|
||||
],
|
||||
)
|
||||
|
||||
stale = "stale-cached-membership"
|
||||
for user_id in ("user-1", "user-2", "user-3"):
|
||||
await cache.async_set_cache(key="team_membership:{}:team-1234".format(user_id), value=stale)
|
||||
await cache.async_set_cache(key="team-1234_{}".format(user_id), value=stale)
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1", "user-2"],
|
||||
update_fields=TeamMemberBulkUpdateFields(allowed_models=["gpt-4"]),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
for user_id in ("user-1", "user-2"):
|
||||
assert await cache.async_get_cache(key="team_membership:{}:team-1234".format(user_id)) is None
|
||||
assert await cache.async_get_cache(key="team-1234_{}".format(user_id)) is None
|
||||
|
||||
# user-3 was never part of this batch; its cache must be untouched.
|
||||
assert await cache.async_get_cache(key="team_membership:user-3:team-1234") == stale
|
||||
assert await cache.async_get_cache(key="team-1234_user-3") == stale
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_role_only_change_does_not_touch_membership_cache(monkeypatch):
|
||||
"""A role-only patch (no budget fields) never writes litellm_teammembership,
|
||||
so it must not evict the membership cache -- there's nothing stale to fix,
|
||||
and clearing it would just cause an avoidable DB read on the next request."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="admin")],
|
||||
)
|
||||
_db, _upsert_mock, _refresh_mock, cache = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
stale = "stale-cached-membership"
|
||||
await cache.async_set_cache(key="team_membership:user-1:team-1234", value=stale)
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(role="user"),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert await cache.async_get_cache(key="team_membership:user-1:team-1234") == stale
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_forwards_shared_default_only_for_members_still_on_it(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
|
|
@ -292,7 +369,7 @@ async def test_bulk_update_forwards_shared_default_only_for_members_still_on_it(
|
|||
],
|
||||
metadata={"team_member_budget_id": "default-bud"},
|
||||
)
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
_db, upsert_mock, _refresh, _cache = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="default-bud")],
|
||||
|
|
@ -325,7 +402,7 @@ async def test_bulk_update_role_writes_team_row_once_and_refreshes_cache(monkeyp
|
|||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
)
|
||||
db, upsert_mock, refresh_mock = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
db, upsert_mock, refresh_mock, _cache = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
|
|
@ -363,7 +440,7 @@ async def test_bulk_update_all_members_in_team_dedups_members(monkeypatch):
|
|||
Member(user_id="user-2", role="user"),
|
||||
],
|
||||
)
|
||||
db, upsert_mock, _refresh = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
db, upsert_mock, _refresh, _cache = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
|
|
@ -386,7 +463,7 @@ async def test_bulk_update_reports_non_members_as_failed(monkeypatch):
|
|||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
_db, upsert_mock, _refresh, _cache = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
|
|
@ -415,7 +492,7 @@ async def test_bulk_update_explicit_null_duration_forwarded_as_patch(monkeypatch
|
|||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
_db, upsert_mock, _refresh, _cache = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue