fix(access groups): sync assigned_key_ids from the key write paths (#36843)
Some checks failed
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-endpoints (push) Waiting to run
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Unit Tests: Core Utilities / core-utils (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Enterprise, Google GenAI & Routing / enterprise-routing (push) Waiting to run
Unit Tests: Integrations (Callbacks & Logging) / integrations (push) Waiting to run
Unit Tests: LLM Provider Transformations / All Other Providers (push) Waiting to run
Unit Tests: LLM Provider Transformations / Vertex AI (push) Waiting to run
Unit Tests: MCP, Secrets, Containers & Misc / misc (push) Waiting to run
Unit Tests: Proxy Auth & Key Management / proxy-auth (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-server (push) Waiting to run
Unit Tests: Proxy Infrastructure / proxy-infra (push) Waiting to run
Unit Tests: Responses, Caching & Types / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Unit Tests: Proxy Legacy Tests / auth-and-jwt (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / key-generation (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-config (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-response-and-misc (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-server (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-server-extras (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-token-counter (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-user-auth-and-spend (push) Has been cancelled
Unit Tests: Proxy Legacy Tests / proxy-utils (push) Has been cancelled

This commit is contained in:
Yassin Kortam 2026-08-13 19:47:02 -07:00 committed by GitHub
parent 86f2f16fd4
commit efbdb6901a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 968 additions and 1 deletions

View file

@ -88,6 +88,11 @@ from litellm.proxy.management_endpoints.common_utils import (
from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_model_to_db,
)
from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_access_group_membership,
sync_key_regeneration_access_group_membership,
sync_key_update_access_group_membership,
)
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
@ -2347,6 +2352,17 @@ async def _process_single_key_update(
proxy_logging_obj=proxy_logging_obj,
)
# After the key's own cache entry is dropped, so a failure here cannot leave the key
# authenticating against the access groups it just lost.
await sync_key_update_access_group_membership(
prisma_client=prisma_client,
key_token=_hash_token_if_needed(
_resolve_token_to_update(data=update_key_request, existing_key_row=existing_key_row)
),
data=update_key_request,
existing_key_row=existing_key_row,
)
# Trigger async hook
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
@ -2828,6 +2844,15 @@ async def update_key_fn(
proxy_logging_obj=proxy_logging_obj,
)
# After the key's own cache entry is dropped, so a failure here cannot leave the key
# authenticating against the access groups it just lost.
await sync_key_update_access_group_membership(
prisma_client=prisma_client,
key_token=_hash_token_if_needed(key),
data=data,
existing_key_row=existing_key_row,
)
if data.spend is not None:
from litellm.proxy.proxy_server import spend_counter_cache
@ -3771,7 +3796,7 @@ async def generate_key_helper_fn(
auto_rotate: bool | None = None,
rotation_interval: str | None = None,
router_settings: dict | None = None,
access_group_ids: list | None = None,
access_group_ids: list[str] | None = None,
budget_limits: list | None = None, # multiple concurrent budget windows
):
from litellm.proxy.proxy_server import premium_user, prisma_client
@ -3979,6 +4004,14 @@ async def generate_key_helper_fn(
create_key_response: Final = await prisma_client.insert_data(data=key_data, table_name="key")
key_data["token_id"] = getattr(create_key_response, "token", None)
created_token_hash: Final = getattr(create_key_response, "token", None)
if isinstance(created_token_hash, str):
await sync_key_access_group_membership(
prisma_client=prisma_client,
key_token=created_token_hash,
previous_access_group_ids=None,
updated_access_group_ids=access_group_ids,
)
key_data["litellm_budget_table"] = getattr(create_key_response, "litellm_budget_table", None)
key_data["created_at"] = getattr(create_key_response, "created_at", None)
key_data["updated_at"] = getattr(create_key_response, "updated_at", None)
@ -4196,6 +4229,7 @@ async def delete_verification_tokens(
deleted_tokens = [key.token for key in authorized_keys]
if len(deleted_tokens) != len(tokens):
failed_tokens = [token for token in tokens if token not in deleted_tokens]
else:
raise Exception("DB not connected. prisma_client is None")
except Exception as e:
@ -4211,6 +4245,16 @@ async def delete_verification_tokens(
hashed_token = hash_token(cast(str, key))
user_api_key_cache.delete_cache(hashed_token)
# After credential invalidation, so a failure here can never keep a deleted key alive.
for deleted_key in authorized_keys:
if deleted_key.token is not None:
await sync_key_access_group_membership(
prisma_client=prisma_client,
key_token=deleted_key.token,
previous_access_group_ids=deleted_key.access_group_ids,
updated_access_group_ids=None,
)
return {
"deleted_keys": deleted_tokens,
"failed_tokens": failed_tokens,
@ -4726,6 +4770,15 @@ async def _execute_virtual_key_regeneration(
proxy_logging_obj=proxy_logging_obj,
)
# After credential invalidation, so a failure here can never keep the old key alive.
await sync_key_regeneration_access_group_membership(
prisma_client=prisma_client,
previous_key_token=hashed_api_key,
new_key_token=new_token_hash,
data=data,
existing_key_row=key_in_db,
)
response: Final = GenerateKeyResponse.model_validate(updated_token_dict)
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(

View file

@ -0,0 +1,173 @@
"""
Reverse sync for the key side of the key <-> access group relationship.
`litellm_accessgrouptable.assigned_key_ids` and `litellm_verificationtoken.access_group_ids`
are the two halves of one relationship and BOTH are read: the access group's
attached-keys view reads the former, and so does the grant check in
`auth_checks.get_authorized_resources_from_key_access_groups`, which authorizes a
key only when the group lists the key's token (or the key's team). The access-group
endpoints maintain both halves already; this module is what the key write paths call
so an edit from that side is mirrored back.
Every write is a single guarded statement rather than a read-modify-write. Prisma has no
atomic scalar-list removal (see `TeamRepository.remove_member`), and the read-modify-write
it otherwise forces is not safe here: a lost update would put an already revoked token back
into a group and restore its grants, or drop a grant an admin just made. The guards also
make each statement idempotent, so a retry cannot duplicate an entry. Each statement covers
every group the request touches at once, so the size of the caller's id list does not turn
into a matching number of round trips, and returns the ids it actually moved so only those
groups are dropped from cache.
It deliberately lives outside `access_group_endpoints`, which is a lazily
registered feature router (see `_lazy_features.LAZY_FEATURES`). Importing that
module eagerly from `key_management_endpoints` would put it in `sys.modules`
without its router ever being included, which drops its routes from the OpenAPI
schema.
"""
from collections.abc import Sequence
from typing import Final, Protocol
from pydantic import BaseModel
from litellm.proxy._types import (
LiteLLM_VerificationToken,
RegenerateKeyRequest,
UpdateKeyRequest,
)
from litellm.proxy.auth.auth_checks import (
_delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive
)
from litellm.repositories.table_repositories import AccessGroupRepository
class _MovedGroupRow(BaseModel):
access_group_id: str
class _RawExecutor(Protocol):
async def query_raw(self, query: str, *args: str | Sequence[str]) -> Sequence[object]: ...
_ATTACH_KEY_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "assigned_key_ids" = array_append("assigned_key_ids", $1) '
'WHERE "access_group_id" = ANY($2::text[]) AND NOT ($1 = ANY("assigned_key_ids")) '
'RETURNING "access_group_id"'
)
_DETACH_KEY_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "assigned_key_ids" = array_remove("assigned_key_ids", $1) '
'WHERE "access_group_id" = ANY($2::text[]) AND $1 = ANY("assigned_key_ids") '
'RETURNING "access_group_id"'
)
_REPOINT_KEY_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "assigned_key_ids" = array_append(array_remove(array_remove("assigned_key_ids", $1), $2), $2) '
'WHERE $1 = ANY("assigned_key_ids") '
'RETURNING "access_group_id"'
)
def _raw_executor(prisma_client: object) -> _RawExecutor:
"""Narrow the untyped Prisma client down to the raw-query call this module makes."""
return AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
async def _invalidate_access_group_cache(access_group_id: str) -> None:
"""
Drop an access group entry from both the in-memory and Redis caches.
Uses a lazy import of user_api_key_cache and proxy_logging_obj from proxy_server
to avoid circular imports, following the same pattern as key_management_endpoints.
"""
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
await _delete_cache_access_object(
access_group_id=access_group_id,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
async def _invalidate_moved_groups(moved_rows: Sequence[object]) -> None:
for row in moved_rows:
await _invalidate_access_group_cache(_MovedGroupRow.model_validate(row).access_group_id)
async def _write_membership(prisma_client: object, sql: str, access_group_ids: frozenset[str], key_token: str) -> None:
"""Run one guarded membership statement for every listed group, dropping the cache of those it moved."""
if not access_group_ids:
return
await _invalidate_moved_groups(
await _raw_executor(prisma_client).query_raw(sql, key_token, sorted(access_group_ids))
)
async def sync_key_access_group_membership(
prisma_client: object,
key_token: str,
previous_access_group_ids: Sequence[str] | None,
updated_access_group_ids: Sequence[str] | None,
) -> None:
"""Mirror a key-side change to `access_group_ids` onto each access group's `assigned_key_ids`."""
previous: Final = frozenset(previous_access_group_ids or ())
updated: Final = frozenset(updated_access_group_ids or ())
await _write_membership(prisma_client, _ATTACH_KEY_SQL, updated - previous, key_token)
await _write_membership(prisma_client, _DETACH_KEY_SQL, previous - updated, key_token)
async def sync_key_update_access_group_membership(
prisma_client: object,
key_token: str,
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
) -> None:
"""
Mirror a key UPDATE onto the group side, honouring `exclude_unset` semantics.
The key row is written from `model_dump(exclude_unset=True)`, so a request that never
mentions `access_group_ids` leaves the key's own list alone and must leave the group's
copy alone too. Reading the attribute instead of `model_fields_set` would see None on
every unrelated edit and withdraw the token from every group it belongs to.
"""
if "access_group_ids" not in data.model_fields_set:
return
await sync_key_access_group_membership(
prisma_client=prisma_client,
key_token=key_token,
previous_access_group_ids=existing_key_row.access_group_ids,
updated_access_group_ids=data.access_group_ids,
)
async def sync_key_regeneration_access_group_membership(
prisma_client: object,
previous_key_token: str,
new_key_token: str,
data: RegenerateKeyRequest | None,
existing_key_row: LiteLLM_VerificationToken,
) -> None:
"""
Re-point every group's copy from the old token to the regenerated one.
Regeneration replaces the token, which is the identity `assigned_key_ids` stores, so
leaving the old hash behind both points the group at a row that no longer exists and
denies the regenerated key the group's grants. The swap is driven by the groups that
hold the old token when the statement runs, not by the key row read earlier, so a group
edited in between is neither resurrected nor skipped. Removing the new token before
appending it keeps a re-run from duplicating it.
"""
await _invalidate_moved_groups(
await _raw_executor(prisma_client).query_raw(_REPOINT_KEY_SQL, previous_key_token, new_key_token)
)
if data is not None:
await sync_key_update_access_group_membership(
prisma_client=prisma_client,
key_token=new_key_token,
data=data,
existing_key_row=existing_key_row,
)

View file

@ -791,6 +791,7 @@ async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch):
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=MagicMock(object_permission_id=None)
)
mock_prisma_client.db.query_raw = AsyncMock(return_value=[])
captured_key_data = {}
@ -15703,3 +15704,743 @@ async def test_key_generate_omitted_budget_duration_still_filled_by_upperbound(m
assert key_row["budget_duration"] == "30d"
assert key_row["budget_reset_at"] is not None
from litellm.proxy.management_helpers.access_group_key_sync import (
_ATTACH_KEY_SQL,
_DETACH_KEY_SQL,
_REPOINT_KEY_SQL,
)
ACCESS_GROUP_SYNC_TOKEN = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
def _access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups):
"""
Back the access group table with an in-memory dict so the sync's writes are observable.
The sync writes through guarded set-based SQL statements, so this emulates exactly what
Postgres does with them, including the guards that make each one idempotent and the
`RETURNING` clause that reports which groups actually moved.
"""
def _repoint(previous_token, new_token):
moved = [
group_id
for group_id, stored in access_groups.items()
if previous_token in stored["assigned_key_ids"]
]
for group_id in moved:
current = access_groups[group_id]["assigned_key_ids"]
access_groups[group_id]["assigned_key_ids"] = [
*(t for t in current if t not in (previous_token, new_token)),
new_token,
]
return moved
def _attach(key_token, access_group_ids):
moved = [
group_id
for group_id in access_group_ids
if group_id in access_groups
and key_token not in access_groups[group_id]["assigned_key_ids"]
]
for group_id in moved:
stored = access_groups[group_id]
stored["assigned_key_ids"] = [*stored["assigned_key_ids"], key_token]
return moved
def _detach(key_token, access_group_ids):
moved = [
group_id
for group_id in access_group_ids
if group_id in access_groups
and key_token in access_groups[group_id]["assigned_key_ids"]
]
for group_id in moved:
stored = access_groups[group_id]
stored["assigned_key_ids"] = [
t for t in stored["assigned_key_ids"] if t != key_token
]
return moved
async def _query_raw(query, *args):
if query == _REPOINT_KEY_SQL:
moved = _repoint(*args)
elif query == _ATTACH_KEY_SQL:
moved = _attach(*args)
else:
assert query == _DETACH_KEY_SQL, f"unexpected statement: {query}"
moved = _detach(*args)
return [{"access_group_id": group_id} for group_id in moved]
raw_mock = AsyncMock(side_effect=_query_raw)
mock_prisma_client.db.query_raw = raw_mock
return raw_mock
async def _authorized_models_for_key(access_groups, token, key_access_group_ids):
"""Run the real auth-time reader against the post-sync access group rows."""
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable
from litellm.proxy.auth.auth_checks import (
get_authorized_resources_from_key_access_groups,
)
async def _get_access_object(*, access_group_id, **_kwargs):
stored = access_groups[access_group_id]
return LiteLLM_AccessGroupTable(
access_group_id=access_group_id,
access_group_name=access_group_id,
access_model_names=list(stored["access_model_names"]),
assigned_team_ids=[],
assigned_key_ids=list(stored["assigned_key_ids"]),
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
side_effect=_get_access_object,
),
):
return await get_authorized_resources_from_key_access_groups(
valid_token=UserAPIKeyAuth(
token=token,
models=[],
team_id="team-a",
access_group_ids=list(key_access_group_ids),
),
team_object=LiteLLM_TeamTable(team_id="team-a", models=[]),
resource_field="access_model_names",
)
@pytest.mark.asyncio
async def test_update_key_syncs_access_group_assigned_key_ids_in_both_directions(
monkeypatch,
):
"""
A key-side edit of `access_group_ids` must be mirrored onto every affected access
group's `assigned_key_ids`, in one operation, in both directions.
`assigned_key_ids` is not display-only. `get_authorized_resources_from_key_access_groups`
reads it as an authorization input and authorizes only when the group lists the key's
token, so a group the key just added must start granting its resources and a group the
key dropped must stop. A single-direction assertion would pass against a fix that only
ever adds (or only ever removes), so this covers add, remove, untouched, and the
authorization consequence of each.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=["ag-drop", "ag-keep"],
)
access_groups = {
"ag-drop": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": ["dropped-model"],
},
"ag-keep": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": ["kept-model"],
},
"ag-add": {"assigned_key_ids": [], "access_model_names": ["added-model"]},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {}})
raw_mock = _access_group_table_mocks(
monkeypatch, mock_prisma_client, access_groups
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
) as invalidate_cache,
):
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(
key=ACCESS_GROUP_SYNC_TOKEN, access_group_ids=["ag-keep", "ag-add"]
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by=None,
)
assert access_groups["ag-drop"]["assigned_key_ids"] == []
assert access_groups["ag-add"]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
assert access_groups["ag-keep"]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
# Both halves go out as single guarded statements. A read-modify-write here lets two
# admins editing one group lose each other's change: an attach can vanish, and a detach
# can put an already revoked token back and restore its grants.
assert sorted(call.args for call in raw_mock.call_args_list) == sorted(
[
(_ATTACH_KEY_SQL, ACCESS_GROUP_SYNC_TOKEN, ["ag-add"]),
(_DETACH_KEY_SQL, ACCESS_GROUP_SYNC_TOKEN, ["ag-drop"]),
]
)
assert {call.args[0] for call in invalidate_cache.call_args_list} == {
"ag-drop",
"ag-add",
}
authorized_models = await _authorized_models_for_key(
access_groups,
ACCESS_GROUP_SYNC_TOKEN,
["ag-drop", "ag-keep", "ag-add"],
)
assert sorted(authorized_models) == ["added-model", "kept-model"]
@pytest.mark.asyncio
async def test_update_key_leaves_access_groups_alone_when_field_is_unset(monkeypatch):
"""
An update that never mentions `access_group_ids` must not touch the group rows.
`prepare_key_update_data` writes from `model_dump(exclude_unset=True)`, so an omitted
field leaves the key row's own list intact. Reading the request attribute instead of
its `model_fields_set` would see None and wipe every group's copy of the token on any
unrelated edit, e.g. a max_budget change.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=["ag-keep"],
)
access_groups = {
"ag-keep": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": ["kept-model"],
},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {}})
raw_mock = _access_group_table_mocks(
monkeypatch, mock_prisma_client, access_groups
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
):
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key=ACCESS_GROUP_SYNC_TOKEN, max_budget=50.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by=None,
)
raw_mock.assert_not_called()
assert access_groups["ag-keep"]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
assert await _authorized_models_for_key(
access_groups, ACCESS_GROUP_SYNC_TOKEN, ["ag-keep"]
) == ["kept-model"]
@pytest.mark.asyncio
async def test_bulk_update_keys_syncs_access_group_assigned_key_ids(monkeypatch):
"""
/key/bulk_update and /team/keys/bulk_update reach the DB through
`_process_single_key_update`, which is a separate write path from /key/update's own
inline one. Both have to maintain the group's copy or a bulk attach grants nothing.
"""
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=["ag-drop"],
)
access_groups = {
"ag-drop": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": ["dropped-model"],
},
"ag-add": {"assigned_key_ids": [], "access_model_names": ["added-model"]},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.update_data = AsyncMock(return_value={"data": {}})
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
new_callable=AsyncMock,
),
):
await _process_single_key_update(
update_key_request=UpdateKeyRequest(
key=ACCESS_GROUP_SYNC_TOKEN, access_group_ids=["ag-add"]
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=AsyncMock(),
proxy_logging_obj=MagicMock(),
llm_router=None,
existing_key_row=key_in_db,
)
assert access_groups["ag-drop"]["assigned_key_ids"] == []
assert access_groups["ag-add"]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
assert await _authorized_models_for_key(
access_groups, ACCESS_GROUP_SYNC_TOKEN, ["ag-drop", "ag-add"]
) == ["added-model"]
@pytest.mark.asyncio
async def test_delete_key_withdraws_token_from_its_access_groups(monkeypatch):
"""
Deleting a key must withdraw its token from every group that lists it.
Without the withdrawal the group keeps a token that no longer resolves to a row, so
the access group page lists a key that does not exist and the list grows without bound.
"""
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=["ag-keep"],
)
access_groups = {
"ag-keep": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN, "other-key"],
"access_model_names": ["kept-model"],
},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[key_in_db]
)
mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_keys": 1})
mock_prisma_client.db.litellm_deletedverificationtoken.create_many = AsyncMock()
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
mock_cache = MagicMock()
mock_cache.delete_cache = MagicMock()
with patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
):
await delete_verification_tokens(
tokens=[ACCESS_GROUP_SYNC_TOKEN],
user_api_key_cache=mock_cache,
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by="admin-user",
)
assert access_groups["ag-keep"]["assigned_key_ids"] == ["other-key"]
@pytest.mark.asyncio
async def test_generate_key_records_token_in_its_access_groups(monkeypatch):
"""
/key/generate with `access_group_ids` must record the new token on the group side.
The key row's own list alone does not authorize: the group has to list the token back
or `get_authorized_resources_from_key_access_groups` contributes nothing, so a key
created against a group silently gets none of its models.
"""
access_groups = {
"ag-add": {"assigned_key_ids": [], "access_model_names": ["added-model"]},
}
created_key = MagicMock()
created_key.token = ACCESS_GROUP_SYNC_TOKEN
created_key.litellm_budget_table = None
created_key.created_at = None
created_key.updated_at = None
mock_prisma_client = AsyncMock()
mock_prisma_client.insert_data = AsyncMock(return_value=created_key)
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.store_audit_logs", False)
with patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
):
await generate_key_helper_fn(
request_type="key",
access_group_ids=["ag-add"],
table_name="key",
user_id="test-user",
)
assert access_groups["ag-add"]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
assert await _authorized_models_for_key(
access_groups, ACCESS_GROUP_SYNC_TOKEN, ["ag-add"]
) == ["added-model"]
@pytest.mark.asyncio
async def test_regenerate_key_repoints_access_group_assigned_key_ids(monkeypatch):
"""
Regeneration replaces the key's token, which is the identity `assigned_key_ids` stores.
Leaving the old hash behind points the group at a token that no longer exists AND
denies the regenerated key the group's grants, so the group's copy has to be
re-pointed from the old hash to the new one in the same operation.
"""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
from litellm.proxy.utils import hash_token
new_token_hash = hash_token("sk-newtoken1234ab12")
existing_key = LiteLLM_VerificationToken(
token="abc123",
user_id="user-1",
models=["gpt-4"],
access_group_ids=["ag-keep"],
)
access_groups = {
"ag-keep": {
"assigned_key_ids": ["abc123"],
"access_model_names": ["kept-model"],
},
}
mock_prisma_client = _make_regenerate_mock_prisma()
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
),
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=existing_key,
hashed_api_key="abc123",
key="abc123",
data=RegenerateKeyRequest(),
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
assert access_groups["ag-keep"]["assigned_key_ids"] == [new_token_hash]
assert await _authorized_models_for_key(
access_groups, new_token_hash, ["ag-keep"]
) == ["kept-model"]
assert (
await _authorized_models_for_key(access_groups, "abc123", ["ag-keep"]) == []
)
@pytest.mark.asyncio
async def test_key_write_paths_revoke_the_key_cache_before_syncing_access_groups(
monkeypatch,
):
"""
Credential invalidation must not sit behind the group sync on any key write path.
The cached auth object still carries the key's old `access_group_ids`, so if the sync
raises first, the request fails with the key still authenticating against groups it
just lost, until that entry expires. Ordering it last means a failed sync degrades to
the stale listing this PR fixes rather than to a stale grant.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
order = []
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=["ag-drop"],
)
access_groups = {
"ag-drop": {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": ["dropped-model"],
},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {}})
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
mock_prisma_client.db.query_raw = AsyncMock(
side_effect=lambda *a, **k: order.append("sync") or []
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
side_effect=lambda **kwargs: order.append("revoke_key_cache"),
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
),
):
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key=ACCESS_GROUP_SYNC_TOKEN, access_group_ids=[]),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by=None,
)
assert order == ["revoke_key_cache", "sync"]
@pytest.mark.asyncio
async def test_update_key_syncs_many_access_groups_in_one_statement_per_direction(
monkeypatch,
):
"""
The number of groups on a request must not become a matching number of round trips.
Anyone allowed to assign access groups picks the size of `access_group_ids`, so a
per-group statement lets one /key/update hold a connection for hundreds of sequential
writes. Both halves are set-based, so the cost is two statements no matter the size.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
dropped = [f"ag-drop-{i}" for i in range(60)]
added = [f"ag-add-{i}" for i in range(60)]
key_in_db = LiteLLM_VerificationToken(
token=ACCESS_GROUP_SYNC_TOKEN,
user_id="test-user",
access_group_ids=dropped,
)
access_groups = {
**{
group_id: {
"assigned_key_ids": [ACCESS_GROUP_SYNC_TOKEN],
"access_model_names": [f"{group_id}-model"],
}
for group_id in dropped
},
**{
group_id: {"assigned_key_ids": [], "access_model_names": [f"{group_id}-model"]}
for group_id in added
},
}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {}})
raw_mock = _access_group_table_mocks(
monkeypatch, mock_prisma_client, access_groups
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
),
):
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(
key=ACCESS_GROUP_SYNC_TOKEN, access_group_ids=added
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
),
litellm_changed_by=None,
)
assert [call.args[0] for call in raw_mock.call_args_list] == [
_ATTACH_KEY_SQL,
_DETACH_KEY_SQL,
]
assert sorted(raw_mock.call_args_list[0].args[2]) == sorted(added)
assert sorted(raw_mock.call_args_list[1].args[2]) == sorted(dropped)
assert all(
access_groups[group_id]["assigned_key_ids"] == [ACCESS_GROUP_SYNC_TOKEN]
for group_id in added
)
assert all(access_groups[group_id]["assigned_key_ids"] == [] for group_id in dropped)
@pytest.mark.asyncio
async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read(
monkeypatch,
):
"""
Regeneration must move whatever the groups hold when it writes, not the key row's list.
That list is read before the new token exists, so replaying it re-adds the key to a
group an admin revoked in between and leaves the dead hash in a group an admin attached
in between, which silently restores one grant and drops another.
"""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
from litellm.proxy.utils import hash_token
new_token_hash = hash_token("sk-newtoken1234ab12")
existing_key = LiteLLM_VerificationToken(
token="abc123",
user_id="user-1",
models=["gpt-4"],
access_group_ids=["ag-revoked-since"],
)
access_groups = {
"ag-revoked-since": {
"assigned_key_ids": [],
"access_model_names": ["revoked-model"],
},
"ag-attached-since": {
"assigned_key_ids": ["abc123"],
"access_model_names": ["attached-model"],
},
}
mock_prisma_client = _make_regenerate_mock_prisma()
_access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_helpers.access_group_key_sync._invalidate_access_group_cache",
new_callable=AsyncMock,
),
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=existing_key,
hashed_api_key="abc123",
key="abc123",
data=RegenerateKeyRequest(),
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
assert access_groups["ag-revoked-since"]["assigned_key_ids"] == []
assert access_groups["ag-attached-since"]["assigned_key_ids"] == [new_token_hash]
assert await _authorized_models_for_key(
access_groups, new_token_hash, ["ag-revoked-since", "ag-attached-since"]
) == ["attached-model"]