fix(auth): read vector store key and team grants through the object permission cache

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-10-04 08:45:53 +00:00
parent 3e5fe50faa
commit 64158253a6
5 changed files with 176 additions and 91 deletions

View file

@ -310,12 +310,6 @@ class _VectorStorePermissionsRow(Protocol):
def vector_stores(self) -> Sequence[str] | None: ...
def _object_permission_table(
repo: _PrismaTableHolder[_VectorStorePermissionsRow],
) -> _PrismaAuthTable[_VectorStorePermissionsRow]:
return _DeadlineBoundedTable(repo.table, "object_permission")
class _PrismaTagRow(Protocol):
tag_name: str
@ -6628,7 +6622,10 @@ def _get_rag_query_vector_store_id(request_body: Mapping[str, object]) -> str |
return vector_store_id if isinstance(vector_store_id, str) and vector_store_id else None
def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool:
GrantLayer = Literal["key", "team", "user"]
def _is_strict_grant_identity(valid_token: UserAPIKeyAuth | None) -> bool:
return (
valid_token is not None
and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS
@ -6636,6 +6633,57 @@ def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool
)
def _strict_grant_layers(
deny_by_default: bool, valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None
) -> frozenset[GrantLayer]:
"""
The identities that must each grant an object under a deny-by-default policy: a virtual key and its team,
a keyless team member's team, or a keyless user's own grant. The master key and dashboard sessions need none
"""
if not deny_by_default or valid_token is None or not _is_strict_grant_identity(valid_token):
return frozenset()
virtual_key: Final = valid_token.via_virtual_key and not valid_token.is_session_token
has_team: Final = team_object is not None or valid_token.team_id is not None
if not virtual_key and not has_team:
return frozenset(("user",))
return frozenset(layer for layer, required in (("key", virtual_key), ("team", has_team)) if required)
async def _identity_grants(
valid_token: UserAPIKeyAuth | None,
team_object: LiteLLM_TeamTable | None,
user_object: LiteLLM_UserTable | None,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> tuple[tuple[GrantLayer, LiteLLM_ObjectPermissionTable | None], ...]:
return (
(
"key",
None
if valid_token is None
else await _cached_object_permission(
valid_token.object_permission_id, valid_token.object_permission, prisma_client, user_api_key_cache
),
),
(
"team",
None
if team_object is None
else await _cached_object_permission(
team_object.object_permission_id, team_object.object_permission, prisma_client, user_api_key_cache
),
),
(
"user",
None
if user_object is None
else await _cached_object_permission(
user_object.object_permission_id, user_object.object_permission, prisma_client, user_api_key_cache
),
),
)
def _vector_store_deny_by_default(general_settings: Mapping[str, object]) -> bool:
"""
Startup rejects a non-boolean value from the config file. A non-boolean value that reaches
@ -6704,7 +6752,7 @@ def _strict_requested_vector_store_ids(request_body: Mapping[str, object]) -> tu
def _require_vector_store_grant(
object_type: Literal["key", "team", "user"],
object_type: GrantLayer,
vector_store_ids_to_run: Sequence[str],
object_permission: _VectorStorePermissionsRow | None,
) -> None:
@ -6722,6 +6770,27 @@ def _require_vector_store_grant(
)
async def _cached_object_permission(
object_permission_id: str | None,
loaded: LiteLLM_ObjectPermissionTable | None,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> LiteLLM_ObjectPermissionTable | None:
"""
The grant row auth already attached to the key, team or user, else the cached row by id. Both are
evicted on every worker when the grant changes, so the request path never reads the table directly.
"""
if object_permission_id is None:
return None
if loaded is not None:
return loaded
return await get_object_permission(
object_permission_id=object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
async def vector_store_access_check(
request_body: dict,
team_object: LiteLLM_TeamTable | None,
@ -6787,56 +6856,23 @@ async def vector_store_access_check(
#########################################################
# Check if the object (key, team, org) has access to the vector store
#########################################################
# Check if the key can access the vector store
key_object_permission: Final = (
await _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique(
where={"object_permission_id": valid_token.object_permission_id},
)
if valid_token is not None and valid_token.object_permission_id is not None
else None
strict_layers: Final = _strict_grant_layers(deny_by_default, valid_token, team_object)
grants: Final = await _identity_grants(
valid_token,
team_object,
user_object if "user" in strict_layers else None,
prisma_client,
user_api_key_cache,
)
strict_identity: Final = deny_by_default and _is_strict_vector_store_identity(valid_token)
strict_key: Final = (
strict_identity and valid_token is not None and valid_token.via_virtual_key and not valid_token.is_session_token
)
has_team: Final = team_object is not None or (valid_token is not None and valid_token.team_id is not None)
if strict_key:
_require_vector_store_grant("key", vector_store_ids_to_run, key_object_permission)
elif key_object_permission is not None:
_can_object_call_vector_stores(
object_type="key",
vector_store_ids_to_run=vector_store_ids_to_run,
object_permissions=key_object_permission,
)
# Check if the team can access the vector store
team_object_permission: Final = (
await _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique(
where={"object_permission_id": team_object.object_permission_id},
)
if team_object is not None and team_object.object_permission_id is not None
else None
)
if strict_identity and has_team:
_require_vector_store_grant("team", vector_store_ids_to_run, team_object_permission)
elif team_object_permission is not None:
_can_object_call_vector_stores(
object_type="team",
vector_store_ids_to_run=vector_store_ids_to_run,
object_permissions=team_object_permission,
)
if strict_identity and not strict_key and not has_team:
user_object_permission: Final = (
await get_object_permission(
object_permission_id=user_object.object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
for layer, grant in grants:
if layer in strict_layers:
_require_vector_store_grant(layer, vector_store_ids_to_run, grant)
elif grant is not None:
_can_object_call_vector_stores(
object_type=layer,
vector_store_ids_to_run=vector_store_ids_to_run,
object_permissions=grant,
)
if user_object is not None and user_object.object_permission_id is not None
else None
)
_require_vector_store_grant("user", vector_store_ids_to_run, user_object_permission)
return True

View file

@ -164,6 +164,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
enforce_all_proxy_mcp_servers_grant_is_admin_only,
handle_update_object_permission_common,
invalidate_cached_object_permissions,
)
from litellm.proxy.management_helpers.team_member_permission_checks import (
TeamMemberPermissionChecks,
@ -2544,6 +2545,10 @@ async def update_team(
verbose_proxy_logger.info("Successfully updated team - %s, info", team_row.team_id)
await sync_team_access_group_membership(prisma_client=prisma_client, team_id=team_row.team_id)
await invalidate_cached_object_permissions(
object_permission_ids=(existing_team.object_permission_id, team_row.object_permission_id),
user_api_key_cache=user_api_key_cache,
)
await _refresh_cached_team(
team_row=team_row,
user_api_key_cache=user_api_key_cache,

View file

@ -57,18 +57,18 @@ def _assert_forbidden_vector_store_denied(gateway: Gateway, model: str, key: str
@pytest.mark.covers("authorization.vector_store.plain_request_skips_object_permission_lookup")
def test_chat_request_without_vector_stores_does_not_read_object_permission_table(gateway: Gateway) -> None:
def test_requests_do_not_read_object_permission_table_once_the_grant_is_cached(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model], object_permission={"vector_stores": ["vs_allowed"]})
_assert_plain_chat_served(gateway, model, key)
before_control: Final = _object_permission_reads()
before_first_use: Final = _object_permission_reads()
_assert_forbidden_vector_store_denied(gateway, model, key)
eventually(_object_permission_reads, lambda reads: reads > before_control, seconds=15)
eventually(_object_permission_reads, lambda reads: reads > before_first_use, seconds=15)
baseline: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic())
for _ in range(PLAIN_REQUESTS):
_assert_plain_chat_served(gateway, model, key)
_assert_forbidden_vector_store_denied(gateway, model, key)
after: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic())
assert after - baseline < PLAIN_REQUESTS, (
f"{PLAIN_REQUESTS} plain chat requests added {after - baseline} object permission reads"
f"{2 * PLAIN_REQUESTS} requests after the first added {after - baseline} object permission reads"
)

View file

@ -1676,8 +1676,9 @@ async def test_vector_store_access_check_with_permissions():
)
mock_prisma_client = MagicMock()
mock_permissions = MagicMock()
mock_permissions.vector_stores = ["store-1", "store-2"]
mock_permissions = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-123", vector_stores=["store-1", "store-2"]
)
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=mock_permissions)
mock_vector_store_registry = MagicMock()
@ -1718,12 +1719,12 @@ async def test_vector_store_access_check_with_team_permissions():
request_body = {}
valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None)
team_object = MagicMock()
team_object.object_permission_id = "team-permission"
team_object = LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission")
mock_prisma_client = MagicMock()
team_permissions = MagicMock()
team_permissions.vector_stores = ["team-store-allowed"]
team_permissions = LiteLLM_ObjectPermissionTable(
object_permission_id="team-permission", vector_stores=["team-store-allowed"]
)
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions)
mock_vector_store_registry = MagicMock()
@ -1785,12 +1786,12 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
}
valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None)
team_object = MagicMock()
team_object.object_permission_id = "team-permission"
team_object = LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission")
mock_prisma_client = MagicMock()
team_permissions = MagicMock()
team_permissions.vector_stores = ["KBALLOWED123"]
team_permissions = LiteLLM_ObjectPermissionTable(
object_permission_id="team-permission", vector_stores=["KBALLOWED123"]
)
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions)
with (
@ -1816,7 +1817,6 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
assert exc_info.value.type == expected_error_type
_KEY_DENIED: Final = ProxyErrorTypes.key_vector_store_access_denied
_TEAM_DENIED: Final = ProxyErrorTypes.team_vector_store_access_denied
_USER_DENIED: Final = ProxyErrorTypes.user_vector_store_access_denied
@ -1839,6 +1839,16 @@ def _virtual_key(
return key
def _permission_row(
object_permission_id: str, permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None
) -> LiteLLM_ObjectPermissionTable | None:
if isinstance(permission, SimpleNamespace):
return LiteLLM_ObjectPermissionTable(
object_permission_id=object_permission_id, vector_stores=permission.vector_stores
)
return permission
async def _common_checks_for_rag_query(
general_settings: Mapping[str, object],
valid_token: UserAPIKeyAuth,
@ -1852,15 +1862,9 @@ async def _common_checks_for_rag_query(
vector_store_registry: VectorStoreRegistry | None = None,
) -> bool:
permissions: Final = {
"key-permission": key_permission,
"team-permission": team_permission,
"user-permission": (
LiteLLM_ObjectPermissionTable(
object_permission_id="user-permission", vector_stores=user_permission.vector_stores
)
if isinstance(user_permission, SimpleNamespace)
else user_permission
),
"key-permission": _permission_row("key-permission", key_permission),
"team-permission": _permission_row("team-permission", team_permission),
"user-permission": _permission_row("user-permission", user_permission),
}
mock_prisma_client: Final = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
@ -1933,7 +1937,11 @@ async def _team_key_rag_query(
@pytest.mark.asyncio
@pytest.mark.parametrize(
("general_settings", "denied_by"),
[({}, None), ({"vector_store_deny_by_default": False}, None), ({"vector_store_deny_by_default": True}, _KEY_DENIED)],
[
({}, None),
({"vector_store_deny_by_default": False}, None),
({"vector_store_deny_by_default": True}, _KEY_DENIED),
],
ids=["flag-omitted", "flag-false", "flag-true"],
)
async def test_standalone_key_without_vector_store_permission_follows_deny_by_default(
@ -1995,7 +2003,9 @@ async def test_proxy_admin_virtual_key_keeps_flag_off_vector_store_behavior(
)
key_permission: Final = None if vector_stores is None else SimpleNamespace(vector_stores=vector_stores)
await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, admin_key, key_permission), denied_by)
await _assert_rag_query_outcome(
_common_checks_for_rag_query(general_settings, admin_key, key_permission), denied_by
)
@pytest.mark.asyncio
@ -2045,7 +2055,9 @@ async def test_team_key_keeps_legacy_vector_store_behavior_when_flag_off(
@pytest.mark.asyncio
async def test_user_owned_team_key_without_key_grant_is_denied_despite_team_membership():
team_member: Final = LiteLLM_UserTable(user_id="key-owner", teams=["team-1"], user_role=LitellmUserRoles.INTERNAL_USER)
team_member: Final = LiteLLM_UserTable(
user_id="key-owner", teams=["team-1"], user_role=LitellmUserRoles.INTERNAL_USER
)
await _assert_rag_query_outcome(
_team_key_rag_query({"vector_store_deny_by_default": True}, [], ["KBSTOREA"], team_member), _KEY_DENIED
@ -2079,8 +2091,7 @@ async def _keyless_rag_query(
),
None if team_vector_stores is None else SimpleNamespace(vector_stores=team_vector_stores),
_user_row(user_vector_stores, teams=() if team_id is None else (team_id,)),
user_permission
or (None if user_vector_stores is None else SimpleNamespace(vector_stores=user_vector_stores)),
user_permission or (None if user_vector_stores is None else SimpleNamespace(vector_stores=user_vector_stores)),
)
@ -2368,6 +2379,30 @@ async def test_keyless_user_grant_is_read_through_the_object_permission_cache():
)
@pytest.mark.asyncio
@pytest.mark.parametrize("deny_by_default", [True, False], ids=["strict", "default"])
async def test_key_and_team_grants_are_read_through_the_object_permission_cache(deny_by_default: bool):
cache: Final = UserApiKeyCache(default_in_memory_ttl=60)
for object_permission_id, vector_stores in (("key-permission", ["KBSTOREA"]), ("team-permission", ["KBSTOREB"])):
await cache.async_set_cache(
key=object_permission_cache_key(object_permission_id),
value=LiteLLM_ObjectPermissionTable(object_permission_id=object_permission_id, vector_stores=vector_stores),
model_type=LiteLLM_ObjectPermissionTable,
)
await _assert_rag_query_outcome(
_common_checks_for_rag_query(
{"vector_store_deny_by_default": deny_by_default},
_virtual_key(object_permission_id="key-permission", team_id="team-1"),
None,
LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission"),
None,
user_api_key_cache=cache,
),
_TEAM_DENIED,
)
def test_can_object_call_model_with_alias():
"""Test that can_object_call_model works with model aliases"""
from litellm import Router
@ -10706,7 +10741,9 @@ async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache
from litellm.proxy._types import LiteLLM_AccessGroupTable
from litellm.proxy.auth.auth_checks import get_access_object
stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"])
stale: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="Policy", access_model_names=["old"]
)
current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []})
client: Final = MagicMock()
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current)
@ -10858,9 +10895,12 @@ async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailab
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
database: Final = MagicMock()
database.get_data = AsyncMock(return_value=UserAPIKeyAuth(
object_permission_id="grant", object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"])
))
database.get_data = AsyncMock(
return_value=UserAPIKeyAuth(
object_permission_id="grant",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"]),
)
)
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(
return_value=None, side_effect=None if missing else RuntimeError("writer unavailable")
)
@ -10883,7 +10923,9 @@ async def test_authoritative_group_grants_propagate_policy_outages(
database: Final = MagicMock()
database.db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(
side_effect=RuntimeError("database unavailable")
)
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
if strict:

View file

@ -8416,6 +8416,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(
# Mock team update
mock_updated_team = MagicMock(spec=LiteLLM_TeamTable)
mock_updated_team.team_id = "org-team-update-bypass-123"
mock_updated_team.object_permission_id = None
mock_updated_team.tpm_limit = 10000
mock_updated_team.rpm_limit = 1000
mock_updated_team.access_group_ids = None
@ -8569,6 +8570,7 @@ async def test_update_team_guardrails_with_org_id(
# Mock team update
mock_updated_team = MagicMock(spec=LiteLLM_TeamTable)
mock_updated_team.team_id = "team-guardrails-123"
mock_updated_team.object_permission_id = None
mock_updated_team.organization_id = "test-org-guardrails"
mock_updated_team.metadata = {
"guardrails": ["aporia-pre-call", "aporia-post-call"]