From 64158253a618c7a9e8fe5c96d46878e46cf9f60b Mon Sep 17 00:00:00 2001 From: mrinal Date: Sun, 4 Oct 2026 08:45:53 +0000 Subject: [PATCH] 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> --- litellm/proxy/auth/auth_checks.py | 148 +++++++++++------- .../management_endpoints/team_endpoints.py | 5 + .../test_object_permission_lookup.py | 10 +- ...st_auth_checks_object_access_and_lookup.py | 102 ++++++++---- .../test_team_endpoints.py | 2 + 5 files changed, 176 insertions(+), 91 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8b42b243cea..539881806c1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8943f5a7416..4389798edae 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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, diff --git a/tests/integration/authorization/test_object_permission_lookup.py b/tests/integration/authorization/test_object_permission_lookup.py index 21fba30c165..a977243e642 100644 --- a/tests/integration/authorization/test_object_permission_lookup.py +++ b/tests/integration/authorization/test_object_permission_lookup.py @@ -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" ) diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index ceb9111a68f..3e3a74d4841 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -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: diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 3e71cc70099..d442c446d81 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -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"]