From f17918c2508ffbd540eb96033d3780c933cb4ead Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 18:39:59 +0000 Subject: [PATCH 01/14] feat(proxy): add opt-in vector_store_deny_by_default for standalone virtual keys Adds general_settings.vector_store_deny_by_default (typed bool, default false). When enabled, a virtual key with no team must list the requested vector store in its object_permission.vector_stores; no permission record, null, or an empty list is denied with key_vector_store_access_denied. Omitted or false keeps the existing behavior, including nonempty allowlist enforcement. The master key is unchanged in both modes. Team keys and keyless callers are deferred to later increments of LIT-6035 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 4 + litellm/proxy/auth/auth_checks.py | 63 ++++++-- ...st_auth_checks_object_access_and_lookup.py | 136 +++++++++++++++++- .../proxy/proxy_server/test_proxy_config.py | 21 ++- 4 files changed, 212 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 988590bfdef..bf95415db23 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2966,6 +2966,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.", ) + vector_store_deny_by_default: bool = Field( + default=False, + description="When True, a virtual key without a team may only use vector stores explicitly listed in its object_permission.vector_stores. A key with no permission record or an empty list is denied. Team keys and non-key callers are not yet covered", + ) missing_session_id: Literal["generate", "reject", "omit"] | None = Field( None, description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400; 'omit' leaves SpendLogs.session_id null, matching callbacks such as Langfuse that only record a client-established metadata.session_id. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index dbd6f28a183..3c7461688eb 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -33,6 +33,7 @@ from litellm.constants import ( DEFAULT_MAX_RECURSE_DEPTH, EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE, END_USER_RESTRICTED_REGISTRY_MAX_SIZE, + LITELLM_PROXY_MASTER_KEY_ALIAS, MODEL_ACCESS_GROUP_REGISTRY_MAX_SIZE, REGISTRY_ERROR_NEGATIVE_CACHE_TTL, TAG_REGISTRY_MAX_SIZE, @@ -44,6 +45,7 @@ from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._types import ( RBAC_ROLES, CallInfo, + ConfigGeneralSettings, LiteLLM_AccessGroupTable, LiteLLM_BudgetTable, LiteLLM_EndUserTable, @@ -1281,6 +1283,15 @@ async def common_checks( request_body=request_body, team_object=team_object, valid_token=valid_token, + deny_by_default=ConfigGeneralSettings.model_validate( + MappingProxyType( + { + "vector_store_deny_by_default": _typed_request_body(general_settings).get( + "vector_store_deny_by_default", False + ) + } + ) + ).vector_store_deny_by_default, ) # 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path) @@ -6622,10 +6633,39 @@ 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_standalone_virtual_key(valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None) -> bool: + return ( + valid_token is not None + and valid_token.via_virtual_key + and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS + and valid_token.team_id is None + and team_object is None + ) + + +def _require_key_vector_store_grant( + vector_store_ids_to_run: Sequence[str], key_object_permission: _VectorStorePermissionsRow | None +) -> None: + if key_object_permission is None or not key_object_permission.vector_stores: + raise ProxyException( + message=f"Key not allowed to access vector store. Tried to access {vector_store_ids_to_run[0]}. vector_store_deny_by_default is enabled and the key has no vector store grants", + type=ProxyErrorTypes.key_vector_store_access_denied, + param="vector_store", + code=status.HTTP_401_UNAUTHORIZED, + ) + _can_object_call_vector_stores( + object_type="key", + vector_store_ids_to_run=vector_store_ids_to_run, + object_permissions=key_object_permission, + ) + + async def vector_store_access_check( request_body: dict, team_object: LiteLLM_TeamTable | None, valid_token: UserAPIKeyAuth | None, + *, + deny_by_default: bool = False, ): """ Checks if the object (key, team, org) has access to the vector store. @@ -6659,18 +6699,21 @@ 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 - if valid_token is not None and valid_token.object_permission_id is not None: - key_object_permission: Final = await _object_permission_table( - ObjectPermissionRepository(prisma_client) - ).find_unique( + key_object_permission: Final = ( + await _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique( where={"object_permission_id": valid_token.object_permission_id}, ) - if 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, - ) + if valid_token is not None and valid_token.object_permission_id is not None + else None + ) + if deny_by_default and _is_standalone_virtual_key(valid_token, team_object): + _require_key_vector_store_grant(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 if team_object is not None and team_object.object_permission_id is not None: 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 6c8b6571991..6d3f30e5958 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 @@ -4,7 +4,7 @@ import json import re import sys import time -from collections.abc import Iterator, Mapping +from collections.abc import Awaitable, Iterator, Mapping from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -61,6 +61,7 @@ from litellm.proxy.auth.auth_checks import ( _check_agent_caller_model_access, _virtual_key_max_budget_check, _virtual_key_soft_budget_check, + common_checks, get_key_object, get_user_object, invalidate_team_member_spend_state, @@ -74,6 +75,7 @@ from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.constants import ( DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, END_USER_RESTRICTED_REGISTRY_MAX_SIZE, + LITELLM_PROXY_MASTER_KEY_ALIAS, PROXY_DB_LOOKUP_MAX_CONCURRENCY, REGISTRY_ERROR_NEGATIVE_CACHE_TTL, TAG_REGISTRY_MAX_SIZE, @@ -1813,6 +1815,138 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query( assert exc_info.value.type == expected_error_type + +def _virtual_key( + object_permission_id: str | None = None, + api_key: str = "sk-standalone", + user_role: LitellmUserRoles | None = None, + team_id: str | None = None, +) -> UserAPIKeyAuth: + key = UserAPIKeyAuth( + api_key=api_key, + user_id="key-owner", + user_role=user_role, + team_id=team_id, + object_permission_id=object_permission_id, + ) + key.via_virtual_key = True + return key + + +async def _common_checks_for_rag_query( + general_settings: Mapping[str, object], + valid_token: UserAPIKeyAuth, + key_permission: SimpleNamespace | None, + team_object: LiteLLM_TeamTable | None = None, +) -> bool: + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=key_permission) + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + request_body = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + } + with ( + patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists + "litellm.proxy.proxy_server.prisma_client", mock_prisma_client + ), + patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists + "litellm.vector_store_registry", None + ), + ): + return await common_checks( + request_body=request_body, + team_object=team_object, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings=dict(general_settings), + route="/v1/rag/query", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=valid_token, + request=MagicMock(spec=Request), + ) + + +async def _assert_rag_query_outcome(request: Awaitable[bool], allowed: bool) -> None: + if allowed: + assert await request is True + return + with pytest.raises(ProxyException) as exc_info: + await request + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( + ProxyErrorTypes.key_vector_store_access_denied, + "vector_store", + "401", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("general_settings", "allowed"), + [({}, True), ({"vector_store_deny_by_default": False}, True), ({"vector_store_deny_by_default": True}, False)], + ids=["flag-omitted", "flag-false", "flag-true"], +) +async def test_standalone_key_without_vector_store_permission_follows_deny_by_default( + general_settings: Mapping[str, object], allowed: bool +): + await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), allowed) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("deny_by_default", "vector_stores", "allowed"), + [ + (True, [], False), + (True, ["KBSTOREA"], True), + (True, ["KBSTOREB"], False), + (True, None, False), + (False, ["KBSTOREB"], False), + ], + ids=["enabled-empty", "enabled-contains", "enabled-excludes", "enabled-null", "disabled-excludes"], +) +async def test_standalone_key_vector_store_permission_record_under_deny_by_default( + deny_by_default: bool, vector_stores: list[str] | None, allowed: bool +): + 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"), + SimpleNamespace(vector_stores=vector_stores), + ), + allowed, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("deny_by_default", [False, True], ids=["flag-false", "flag-true"]) +async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_default: bool): + master_key = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN) + + await _assert_rag_query_outcome( + _common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), True + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("valid_token", "team_object", "allowed"), + [ + (_virtual_key(user_role=LitellmUserRoles.PROXY_ADMIN), None, False), + (_virtual_key(team_id="team-1"), LiteLLM_TeamTable(team_id="team-1"), True), + ], + ids=["admin-owned-standalone-key-denied", "team-key-deferred"], +) +async def test_deny_by_default_scope_is_standalone_virtual_keys( + valid_token: UserAPIKeyAuth, team_object: LiteLLM_TeamTable | None, allowed: bool +): + await _assert_rag_query_outcome( + _common_checks_for_rag_query({"vector_store_deny_by_default": True}, valid_token, None, team_object), allowed + ) + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 7fdf9277154..0c7a4a1608b 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -26,7 +26,7 @@ import pytest from pydantic import JsonValue, TypeAdapter, ValidationError import litellm -from litellm.proxy._types import CommonProxyErrors +from litellm.proxy._types import CommonProxyErrors, ConfigGeneralSettings from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.proxy.proxy_server import ( ProxyConfig, @@ -2274,6 +2274,25 @@ async def test_load_config_logs_disabled_budget_reservation_once(tmp_path, monke assert [record.levelno for record in records] == ([logging.INFO] if setting == "true" else []) +@pytest.mark.asyncio +@pytest.mark.parametrize(("yaml_value", "expected"), [("true", True), ("false", False)]) +async def test_load_config_yaml_vector_store_deny_by_default_is_boolean( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str, expected: bool +): + config_file = tmp_path / "vector_store.yaml" + config_file.write_text( + f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n" + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + _, _, general_settings = await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + + assert general_settings["vector_store_deny_by_default"] is expected + assert ConfigGeneralSettings.model_validate(dict(general_settings)).vector_store_deny_by_default is expected + + @pytest.mark.asyncio async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch): """Regression: router_settings.plugins dotted-path strings must be resolved to From 5a892fc548f25e91da49f5232ec3c29604e4f68b Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 19:05:27 +0000 Subject: [PATCH 02/14] feat(proxy): require key and team vector store grants for team keys under vector_store_deny_by_default With the flag enabled, a virtual key on a team needs both its own grant and its team's grant for every requested vector store. A missing permission record, an empty list or an unresolved team grants nothing. Dashboard session keys and the master key keep their existing behavior, and flag-off behavior is unchanged Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 2 +- litellm/proxy/auth/auth_checks.py | 50 ++++--- ...st_auth_checks_object_access_and_lookup.py | 137 ++++++++++++++---- .../test_user_api_key_auth_request_flow.py | 62 ++++++++ 4 files changed, 200 insertions(+), 51 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bf95415db23..4526765edcf 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2968,7 +2968,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) vector_store_deny_by_default: bool = Field( default=False, - description="When True, a virtual key without a team may only use vector stores explicitly listed in its object_permission.vector_stores. A key with no permission record or an empty list is denied. Team keys and non-key callers are not yet covered", + description="When True, a virtual key may only use vector stores explicitly listed in its object_permission.vector_stores, and a key on a team also needs the team to list them. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys and non-key callers are not yet covered", ) missing_session_id: Literal["generate", "reject", "omit"] | None = Field( None, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3c7461688eb..04003dc8187 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -44,6 +44,7 @@ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._types import ( RBAC_ROLES, + UI_TEAM_ID, CallInfo, ConfigGeneralSettings, LiteLLM_AccessGroupTable, @@ -6633,30 +6634,31 @@ 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_standalone_virtual_key(valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None) -> bool: +def _is_strict_vector_store_virtual_key(valid_token: UserAPIKeyAuth | None) -> bool: return ( valid_token is not None and valid_token.via_virtual_key and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS - and valid_token.team_id is None - and team_object is None + and valid_token.team_id != UI_TEAM_ID ) -def _require_key_vector_store_grant( - vector_store_ids_to_run: Sequence[str], key_object_permission: _VectorStorePermissionsRow | None +def _require_vector_store_grant( + object_type: Literal["key", "team"], + vector_store_ids_to_run: Sequence[str], + object_permission: _VectorStorePermissionsRow | None, ) -> None: - if key_object_permission is None or not key_object_permission.vector_stores: + if object_permission is None or not object_permission.vector_stores: raise ProxyException( - message=f"Key not allowed to access vector store. Tried to access {vector_store_ids_to_run[0]}. vector_store_deny_by_default is enabled and the key has no vector store grants", - type=ProxyErrorTypes.key_vector_store_access_denied, + message=f"{object_type.capitalize()} not allowed to access vector store. Tried to access {vector_store_ids_to_run[0]}. vector_store_deny_by_default is enabled and the {object_type} has no vector store grants", + type=ProxyErrorTypes.get_vector_store_access_error_type_for_object(object_type), param="vector_store", code=status.HTTP_401_UNAUTHORIZED, ) _can_object_call_vector_stores( - object_type="key", + object_type=object_type, vector_store_ids_to_run=vector_store_ids_to_run, - object_permissions=key_object_permission, + object_permissions=object_permission, ) @@ -6706,8 +6708,9 @@ async def vector_store_access_check( if valid_token is not None and valid_token.object_permission_id is not None else None ) - if deny_by_default and _is_standalone_virtual_key(valid_token, team_object): - _require_key_vector_store_grant(vector_store_ids_to_run, key_object_permission) + strict_key: Final = deny_by_default and _is_strict_vector_store_virtual_key(valid_token) + 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", @@ -6716,18 +6719,21 @@ async def vector_store_access_check( ) # Check if the team can access the vector store - if team_object is not None and team_object.object_permission_id is not None: - team_object_permission: Final = await _object_permission_table( - ObjectPermissionRepository(prisma_client) - ).find_unique( + 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_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 team_object is not None and team_object.object_permission_id is not None + else None + ) + if strict_key and (team_object is not None or (valid_token is not None and valid_token.team_id is not None)): + _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, + ) return True 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 6d3f30e5958..f7e58612b84 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 @@ -1816,6 +1816,10 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query( +_KEY_DENIED: Final = ProxyErrorTypes.key_vector_store_access_denied +_TEAM_DENIED: Final = ProxyErrorTypes.team_vector_store_access_denied + + def _virtual_key( object_permission_id: str | None = None, api_key: str = "sk-standalone", @@ -1838,9 +1842,14 @@ async def _common_checks_for_rag_query( valid_token: UserAPIKeyAuth, key_permission: SimpleNamespace | None, team_object: LiteLLM_TeamTable | None = None, + team_permission: SimpleNamespace | None = None, + user_object: LiteLLM_UserTable | None = None, ) -> bool: + permissions: Final = {"key-permission": key_permission, "team-permission": team_permission} mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=key_permission) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + side_effect=lambda where: permissions.get(where["object_permission_id"]) + ) mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) request_body = { "model": "gpt-4o-mini", @@ -1858,7 +1867,7 @@ async def _common_checks_for_rag_query( return await common_checks( request_body=request_body, team_object=team_object, - user_object=None, + user_object=user_object, end_user_object=None, global_proxy_spend=None, general_settings=dict(general_settings), @@ -1870,45 +1879,59 @@ async def _common_checks_for_rag_query( ) -async def _assert_rag_query_outcome(request: Awaitable[bool], allowed: bool) -> None: - if allowed: +async def _assert_rag_query_outcome(request: Awaitable[bool], denied_by: ProxyErrorTypes | None) -> None: + if denied_by is None: assert await request is True return with pytest.raises(ProxyException) as exc_info: await request - assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( - ProxyErrorTypes.key_vector_store_access_denied, - "vector_store", - "401", + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (denied_by, "vector_store", "401") + + +async def _team_key_rag_query( + general_settings: Mapping[str, object], + key_vector_stores: list[str] | None, + team_vector_stores: list[str] | None, + user_object: LiteLLM_UserTable | None = None, +) -> bool: + return await _common_checks_for_rag_query( + general_settings, + _virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission", team_id="team-1"), + None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores), + LiteLLM_TeamTable( + team_id="team-1", object_permission_id=None if team_vector_stores is None else "team-permission" + ), + None if team_vector_stores is None else SimpleNamespace(vector_stores=team_vector_stores), + user_object, ) @pytest.mark.asyncio @pytest.mark.parametrize( - ("general_settings", "allowed"), - [({}, True), ({"vector_store_deny_by_default": False}, True), ({"vector_store_deny_by_default": True}, False)], + ("general_settings", "denied_by"), + [({}, 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( - general_settings: Mapping[str, object], allowed: bool + general_settings: Mapping[str, object], denied_by: ProxyErrorTypes | None ): - await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), allowed) + await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), denied_by) @pytest.mark.asyncio @pytest.mark.parametrize( - ("deny_by_default", "vector_stores", "allowed"), + ("deny_by_default", "vector_stores", "denied_by"), [ - (True, [], False), - (True, ["KBSTOREA"], True), - (True, ["KBSTOREB"], False), - (True, None, False), - (False, ["KBSTOREB"], False), + (True, [], _KEY_DENIED), + (True, ["KBSTOREA"], None), + (True, ["KBSTOREB"], _KEY_DENIED), + (True, None, _KEY_DENIED), + (False, ["KBSTOREB"], _KEY_DENIED), ], ids=["enabled-empty", "enabled-contains", "enabled-excludes", "enabled-null", "disabled-excludes"], ) async def test_standalone_key_vector_store_permission_record_under_deny_by_default( - deny_by_default: bool, vector_stores: list[str] | None, allowed: bool + deny_by_default: bool, vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None ): await _assert_rag_query_outcome( _common_checks_for_rag_query( @@ -1916,7 +1939,7 @@ async def test_standalone_key_vector_store_permission_record_under_deny_by_defau _virtual_key(object_permission_id="key-permission"), SimpleNamespace(vector_stores=vector_stores), ), - allowed, + denied_by, ) @@ -1926,24 +1949,82 @@ async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_defa master_key = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN) await _assert_rag_query_outcome( - _common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), True + _common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), None ) @pytest.mark.asyncio @pytest.mark.parametrize( - ("valid_token", "team_object", "allowed"), + ("general_settings", "vector_stores", "denied_by"), [ - (_virtual_key(user_role=LitellmUserRoles.PROXY_ADMIN), None, False), - (_virtual_key(team_id="team-1"), LiteLLM_TeamTable(team_id="team-1"), True), + ({}, None, None), + ({"vector_store_deny_by_default": False}, None, None), + ({"vector_store_deny_by_default": False}, ["KBSTOREB"], _KEY_DENIED), ], - ids=["admin-owned-standalone-key-denied", "team-key-deferred"], + ids=["flag-omitted-no-record", "flag-false-no-record", "flag-false-excludes"], ) -async def test_deny_by_default_scope_is_standalone_virtual_keys( - valid_token: UserAPIKeyAuth, team_object: LiteLLM_TeamTable | None, allowed: bool +async def test_proxy_admin_virtual_key_keeps_flag_off_vector_store_behavior( + general_settings: Mapping[str, object], vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None +): + admin_key = _virtual_key( + object_permission_id=None if vector_stores is None else "key-permission", user_role=LitellmUserRoles.PROXY_ADMIN + ) + key_permission = 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) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("key_vector_stores", "team_vector_stores", "denied_by"), + [ + (["KBSTOREA"], ["KBSTOREA"], None), + (None, ["KBSTOREA"], _KEY_DENIED), + ([], ["KBSTOREA"], _KEY_DENIED), + (["KBSTOREA"], None, _TEAM_DENIED), + (["KBSTOREA"], [], _TEAM_DENIED), + (["KBSTOREB"], ["KBSTOREA"], _KEY_DENIED), + (["KBSTOREA"], ["KBSTOREB"], _TEAM_DENIED), + ], + ids=["both-grant", "key-no-record", "key-empty", "team-no-record", "team-empty", "key-excludes", "team-excludes"], +) +async def test_team_key_requires_key_and_team_grant_under_deny_by_default( + key_vector_stores: list[str] | None, team_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None ): await _assert_rag_query_outcome( - _common_checks_for_rag_query({"vector_store_deny_by_default": True}, valid_token, None, team_object), allowed + _team_key_rag_query({"vector_store_deny_by_default": True}, key_vector_stores, team_vector_stores), denied_by + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("general_settings", "key_vector_stores", "team_vector_stores", "denied_by"), + [ + ({}, [], ["KBSTOREA"], None), + ({"vector_store_deny_by_default": False}, [], ["KBSTOREA"], None), + ({"vector_store_deny_by_default": False}, None, None, None), + ({"vector_store_deny_by_default": False}, ["KBSTOREB"], ["KBSTOREA"], _KEY_DENIED), + ({"vector_store_deny_by_default": False}, ["KBSTOREA"], ["KBSTOREB"], _TEAM_DENIED), + ], + ids=["omitted-key-empty", "false-key-empty", "false-no-records", "false-key-excludes", "false-team-excludes"], +) +async def test_team_key_keeps_legacy_vector_store_behavior_when_flag_off( + general_settings: Mapping[str, object], + key_vector_stores: list[str] | None, + team_vector_stores: list[str] | None, + denied_by: ProxyErrorTypes | None, +): + await _assert_rag_query_outcome( + _team_key_rag_query(general_settings, key_vector_stores, team_vector_stores), denied_by + ) + + +@pytest.mark.asyncio +async def test_user_owned_team_key_without_key_grant_is_denied_despite_team_membership(): + team_member = 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 ) diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 781d0a13bfd..2fe435e733e 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -4854,6 +4854,68 @@ async def test_centralized_common_checks_tolerates_db_errors_when_fetching_conte setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("deny_by_default", "team_lookup_error", "denied"), + [ + (False, RuntimeError("team cache unavailable"), False), + (True, RuntimeError("team cache unavailable"), True), + (True, HTTPException(status_code=404, detail="team read failed"), True), + ], + ids=["flag-off-lookup-swallowed", "flag-on-lookup-swallowed", "flag-on-team-rebuilt-from-token"], +) +async def test_team_key_vector_store_access_when_team_cannot_be_resolved( + monkeypatch: pytest.MonkeyPatch, deny_by_default: bool, team_lookup_error: Exception, denied: bool +): + from fastapi import Request + from starlette.datastructures import URL + + token = UserAPIKeyAuth( + api_key="sk-team-key", team_id="team-1", team_models=["gpt-4o-mini"], object_permission_id="key-permission" + ) + token.via_virtual_key = True + request = Request(scope={"type": "http"}) + request._url = URL(url="/v1/rag/query") + database = MagicMock() + database.db.litellm_objectpermissiontable.find_unique = AsyncMock( + side_effect=lambda where: SimpleNamespace(vector_stores=["KBSTOREA"]) + if where["object_permission_id"] == "key-permission" + else None + ) + attrs = { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "general_settings": {"vector_store_deny_by_default": deny_by_default}, + } + for name, value in attrs.items(): + monkeypatch.setattr(litellm.proxy.proxy_server, name, value) + monkeypatch.setattr(litellm, "vector_store_registry", None) + monkeypatch.setattr( + "litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error) + ) + + checks = _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + }, + route="/v1/rag/query", + ) + if not denied: + await checks + return + with pytest.raises(ProxyException) as exc_info: + await checks + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( + ProxyErrorTypes.team_vector_store_access_denied, + "vector_store", + "401", + ) + + @pytest.mark.asyncio async def test_centralized_common_checks_propagates_end_user_budget_error(): """Regression: ``get_end_user_object`` raises ``litellm.BudgetExceededError`` From da613937d53d84fdd318db2097e9fa0a99a70d05 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 19:36:15 +0000 Subject: [PATCH 03/14] feat(proxy): require user or team vector store grants for keyless requests under vector_store_deny_by_default Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_core_utils/error_normalization.py | 1 + litellm/proxy/_types.py | 11 +- litellm/proxy/auth/auth_checks.py | 27 +++- ...st_auth_checks_object_access_and_lookup.py | 145 +++++++++++++++++- .../test_user_api_key_auth_request_flow.py | 81 ++++++++++ 5 files changed, 256 insertions(+), 9 deletions(-) diff --git a/litellm/litellm_core_utils/error_normalization.py b/litellm/litellm_core_utils/error_normalization.py index be0098ec34b..afb7a0e9b65 100644 --- a/litellm/litellm_core_utils/error_normalization.py +++ b/litellm/litellm_core_utils/error_normalization.py @@ -106,6 +106,7 @@ _PROXY_ERROR_TYPE_MAP: Final[Mapping[str, str]] = MappingProxyType( "key_vector_store_access_denied": PERMISSION_DENIED, "team_vector_store_access_denied": PERMISSION_DENIED, "org_vector_store_access_denied": PERMISSION_DENIED, + "user_vector_store_access_denied": PERMISSION_DENIED, "tool_access_denied": PERMISSION_DENIED, "team_member_permission_error": PERMISSION_DENIED, "not_found_error": RESOURCE_NOT_FOUND, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 4526765edcf..fe505ad6b06 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2968,7 +2968,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) vector_store_deny_by_default: bool = Field( default=False, - description="When True, a virtual key may only use vector stores explicitly listed in its object_permission.vector_stores, and a key on a team also needs the team to list them. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys and non-key callers are not yet covered", + description="When True, a vector store must be explicitly listed in object_permission.vector_stores: a virtual key needs its own grant plus its team's, a keyless team member needs the team's, and a user with neither needs their own. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys are not yet covered", ) missing_session_id: Literal["generate", "reject", "omit"] | None = Field( None, @@ -4514,6 +4514,11 @@ class ProxyErrorTypes(str, enum.Enum): Organization does not have access to the vector store """ + user_vector_store_access_denied = "user_vector_store_access_denied" + """ + User does not have access to the vector store + """ + team_member_already_in_team = "team_member_already_in_team" """ Team member is already in team @@ -4546,7 +4551,7 @@ class ProxyErrorTypes(str, enum.Enum): @classmethod def get_vector_store_access_error_type_for_object( - cls, object_type: Literal["key", "team", "org"] + cls, object_type: Literal["key", "team", "org", "user"] ) -> "ProxyErrorTypes": """ Get the vector store access error type for object_type @@ -4557,6 +4562,8 @@ class ProxyErrorTypes(str, enum.Enum): return cls.team_vector_store_access_denied elif object_type == "org": return cls.org_vector_store_access_denied + elif object_type == "user": + return cls.user_vector_store_access_denied DB_CONNECTION_ERROR_TYPES: Final = ( diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 04003dc8187..9a4f72441cb 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1284,6 +1284,7 @@ async def common_checks( request_body=request_body, team_object=team_object, valid_token=valid_token, + user_object=user_object, deny_by_default=ConfigGeneralSettings.model_validate( MappingProxyType( { @@ -6634,17 +6635,16 @@ 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_virtual_key(valid_token: UserAPIKeyAuth | None) -> bool: +def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool: return ( valid_token is not None - and valid_token.via_virtual_key and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS and valid_token.team_id != UI_TEAM_ID ) def _require_vector_store_grant( - object_type: Literal["key", "team"], + object_type: Literal["key", "team", "user"], vector_store_ids_to_run: Sequence[str], object_permission: _VectorStorePermissionsRow | None, ) -> None: @@ -6667,6 +6667,7 @@ async def vector_store_access_check( team_object: LiteLLM_TeamTable | None, valid_token: UserAPIKeyAuth | None, *, + user_object: LiteLLM_UserTable | None = None, deny_by_default: bool = False, ): """ @@ -6708,7 +6709,11 @@ async def vector_store_access_check( if valid_token is not None and valid_token.object_permission_id is not None else None ) - strict_key: Final = deny_by_default and _is_strict_vector_store_virtual_key(valid_token) + 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: @@ -6726,7 +6731,7 @@ async def vector_store_access_check( if team_object is not None and team_object.object_permission_id is not None else None ) - if strict_key and (team_object is not None or (valid_token is not None and valid_token.team_id is not 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( @@ -6734,11 +6739,21 @@ async def vector_store_access_check( 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 _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique( + where={"object_permission_id": user_object.object_permission_id}, + ) + 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 def _can_object_call_vector_stores( - object_type: Literal["key", "team", "org"], + object_type: Literal["key", "team", "org", "user"], vector_store_ids_to_run: Sequence[str], object_permissions: _VectorStorePermissionsRow | None, ): 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 f7e58612b84..436a773b85f 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 @@ -1818,6 +1818,7 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query( _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 def _virtual_key( @@ -1844,8 +1845,13 @@ async def _common_checks_for_rag_query( team_object: LiteLLM_TeamTable | None = None, team_permission: SimpleNamespace | None = None, user_object: LiteLLM_UserTable | None = None, + user_permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None = None, ) -> bool: - permissions: Final = {"key-permission": key_permission, "team-permission": team_permission} + permissions: Final = { + "key-permission": key_permission, + "team-permission": team_permission, + "user-permission": user_permission, + } mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( side_effect=lambda where: permissions.get(where["object_permission_id"]) @@ -2028,6 +2034,143 @@ async def test_user_owned_team_key_without_key_grant_is_denied_despite_team_memb ) +def _user_row(user_vector_stores: list[str] | None, teams: tuple[str, ...] = ()) -> LiteLLM_UserTable: + return LiteLLM_UserTable( + user_id="user-1", + teams=list(teams), + user_role=LitellmUserRoles.INTERNAL_USER, + object_permission_id=None if user_vector_stores is None else "user-permission", + ) + + +async def _keyless_rag_query( + general_settings: Mapping[str, object], + user_vector_stores: list[str] | None, + team_vector_stores: list[str] | None = None, + team_id: str | None = None, + user_permission: LiteLLM_ObjectPermissionTable | None = None, +) -> bool: + return await _common_checks_for_rag_query( + general_settings, + UserAPIKeyAuth(user_id="user-1", team_id=team_id, user_role=LitellmUserRoles.INTERNAL_USER), + None, + None + if team_id is None + else LiteLLM_TeamTable( + team_id=team_id, object_permission_id=None if team_vector_stores is None else "team-permission" + ), + 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)), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_vector_stores", "denied_by"), + [(["KBSTOREA"], None), (["KBSTOREB"], _USER_DENIED), ([], _USER_DENIED), (None, _USER_DENIED)], + ids=["user-grants", "user-excludes", "user-empty", "user-no-record"], +) +async def test_keyless_user_without_team_needs_user_grant_under_deny_by_default( + user_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None +): + await _assert_rag_query_outcome( + _keyless_rag_query({"vector_store_deny_by_default": True}, user_vector_stores), denied_by + ) + + +@pytest.mark.asyncio +async def test_keyless_user_null_vector_store_list_grants_nothing_under_deny_by_default(): + null_list: Final = LiteLLM_ObjectPermissionTable(object_permission_id="user-permission", vector_stores=None) + + await _assert_rag_query_outcome( + _keyless_rag_query({"vector_store_deny_by_default": True}, [], user_permission=null_list), _USER_DENIED + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("general_settings", "user_vector_stores"), + [({}, ["KBSTOREB"]), ({"vector_store_deny_by_default": False}, ["KBSTOREB"]), ({}, None)], + ids=["omitted-user-excludes", "false-user-excludes", "omitted-user-no-record"], +) +async def test_keyless_user_keeps_legacy_vector_store_behavior_when_flag_off( + general_settings: Mapping[str, object], user_vector_stores: list[str] | None +): + await _assert_rag_query_outcome(_keyless_rag_query(general_settings, user_vector_stores), None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("team_vector_stores", "user_vector_stores", "denied_by"), + [ + (["KBSTOREA"], [], None), + (["KBSTOREB"], ["KBSTOREA"], _TEAM_DENIED), + ([], ["KBSTOREA"], _TEAM_DENIED), + (None, ["KBSTOREA"], _TEAM_DENIED), + ], + ids=["team-grants-user-empty", "team-excludes-user-grants", "team-empty-user-grants", "team-no-record-user-grants"], +) +async def test_keyless_team_member_needs_only_team_grant_under_deny_by_default( + team_vector_stores: list[str] | None, user_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None +): + await _assert_rag_query_outcome( + _keyless_rag_query({"vector_store_deny_by_default": True}, user_vector_stores, team_vector_stores, "team-1"), + denied_by, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("team_id", "denied_by"), + [("team-1", None), (None, None)], + ids=["session-on-team-uses-team-grant", "session-without-team-uses-user-grant"], +) +async def test_session_token_is_not_a_virtual_key_under_deny_by_default( + team_id: str | None, denied_by: ProxyErrorTypes | None +): + session = UserAPIKeyAuth( + api_key="hashed-session-token", + user_id="user-1", + team_id=team_id, + user_role=LitellmUserRoles.INTERNAL_USER, + is_session_token=True, + ) + session.via_virtual_key = True + + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + session, + None, + None if team_id is None else LiteLLM_TeamTable(team_id=team_id, object_permission_id="team-permission"), + SimpleNamespace(vector_stores=["KBSTOREA"]), + _user_row(["KBSTOREA"] if team_id is None else []), + SimpleNamespace(vector_stores=["KBSTOREA"] if team_id is None else []), + ), + denied_by, + ) + + +@pytest.mark.asyncio +async def test_user_owned_standalone_key_cannot_use_owner_grants_under_deny_by_default(): + owner: Final = LiteLLM_UserTable( + user_id="key-owner", user_role=LitellmUserRoles.INTERNAL_USER, object_permission_id="user-permission" + ) + + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + _virtual_key(), + None, + user_object=owner, + user_permission=SimpleNamespace(vector_stores=["KBSTOREA"]), + ), + _KEY_DENIED, + ) + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 2fe435e733e..90078fa836f 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -4916,6 +4916,87 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("team_lookup", "denied"), + [ + ({"team-1": "team-1-grants-a"}, False), + (RuntimeError("team cache unavailable"), True), + ({"team-1": "team-1-grants-none", "team-2": "team-2-grants-a"}, True), + ], + ids=["resolved-team-grants", "team-lookup-swallowed-no-personal-fallback", "other-member-team-grants-ignored"], +) +async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_team( + monkeypatch: pytest.MonkeyPatch, team_lookup: dict[str, str] | Exception, denied: bool +): + from fastapi import Request + from starlette.datastructures import URL + + token = UserAPIKeyAuth( + user_id="user-1", + team_id="team-1", + team_models=["gpt-4o-mini"], + user_role=LitellmUserRoles.INTERNAL_USER, + ) + request = Request(scope={"type": "http"}) + request._url = URL(url="/v1/rag/query") + grants = { + "team-1-grants-a": ["KBSTOREA"], + "team-1-grants-none": [], + "team-2-grants-a": ["KBSTOREA"], + "user-grants-a": ["KBSTOREA"], + } + database = MagicMock() + database.db.litellm_objectpermissiontable.find_unique = AsyncMock( + side_effect=lambda where: SimpleNamespace(vector_stores=grants[where["object_permission_id"]]) + ) + attrs = { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "general_settings": {"vector_store_deny_by_default": True}, + } + for name, value in attrs.items(): + monkeypatch.setattr(litellm.proxy.proxy_server, name, value) + monkeypatch.setattr(litellm, "vector_store_registry", None) + + async def get_team(team_id: str, **_: object) -> LiteLLM_TeamTableCachedObj: + if isinstance(team_lookup, Exception): + raise team_lookup + return LiteLLM_TeamTableCachedObj(team_id=team_id, models=["gpt-4o-mini"], object_permission_id=team_lookup[team_id]) + + monkeypatch.setattr("litellm.proxy.auth.user_api_key_auth.get_team_object", get_team) + monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_membership", AsyncMock(return_value=None)) + monkeypatch.setattr( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + AsyncMock( + return_value=LiteLLM_UserTable( + user_id="user-1", teams=["team-1", "team-2"], object_permission_id="user-grants-a" + ) + ), + ) + + checks = _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + }, + route="/v1/rag/query", + ) + if not denied: + await checks + return + with pytest.raises(ProxyException) as exc_info: + await checks + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( + ProxyErrorTypes.team_vector_store_access_denied, + "vector_store", + "401", + ) + + @pytest.mark.asyncio async def test_centralized_common_checks_propagates_end_user_budget_error(): """Regression: ``get_end_user_object`` raises ``litellm.BudgetExceededError`` From d3c21f716e5bdeffa14c088564a6bf667e91ad42 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 20:19:10 +0000 Subject: [PATCH 04/14] test(proxy): cover vector_store_deny_by_default through a real proxy with key, team and JWT identities Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_rag_query_vector_store_allowlist.py | 237 ++++++++++++++++++ ...st_auth_checks_object_access_and_lookup.py | 71 +++++- 2 files changed, 301 insertions(+), 7 deletions(-) diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index 896c88c68bb..4f1270a33aa 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -1,5 +1,7 @@ from __future__ import annotations +import json +import time import uuid from collections.abc import Iterator, Mapping from pathlib import Path @@ -7,16 +9,20 @@ from types import MappingProxyType from typing import Final, Literal, TypeAlias import httpx +import jwt import pytest import yaml +from cryptography.hazmat.primitives.asymmetric import rsa from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server from integration.authorization._guardrail_opt_out import upstream_observations from pydantic import JsonValue CONFIG_STORE_ID: Final = "vs_integration_config_store" PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",) +JWT_KEY_ID: Final = "integration-vector-store-jwt-key" JsonObject: TypeAlias = dict[str, JsonValue] @@ -99,6 +105,209 @@ def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[t yield no_registry_gateway, upstream_gateway +class StrictGateway: + def __init__(self, gateway: Gateway, upstream: Gateway, signing_key: rsa.RSAPrivateKey) -> None: + self.gateway: Final = gateway + self.upstream: Final = upstream + self._signing_key: Final = signing_key + + def jwt(self, subject: str, groups: tuple[str, ...] = ()) -> str: + claims: Final[JsonObject] = { + "sub": subject, + "groups": _json_array(*groups), + "iat": int(time.time()), + "exp": int(time.time()) + 300, + } + return jwt.encode(claims, self._signing_key, algorithm="RS256", headers={"kid": JWT_KEY_ID}) + + +def _strict_config(directory: Path, *, deny_by_default: bool = True) -> Path: + config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text())) + general_settings: Final = object_value(config["general_settings"]) + strict: Final[JsonObject] = { + **config, + "general_settings": { + **general_settings, + "vector_store_deny_by_default": deny_by_default, + "enable_jwt_auth": True, + "litellm_jwtauth": { + "user_id_jwt_field": "sub", + "team_ids_jwt_field": "groups", + "team_allowed_routes": _json_array("openai_routes", "info_routes", "/v1/rag/query"), + }, + }, + } + path: Final = directory / f"proxy_vector_store_deny_by_default_{deny_by_default}.yaml" + path.write_text(yaml.safe_dump(strict)) + return path + + +@pytest.fixture(scope="module") +def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictGateway]: + signing_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(signing_key.public_key()) + jwks_body: Final = json.dumps({"keys": [{**json.loads(public_jwk), "kid": JWT_KEY_ID}]}).encode() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return Reply(body=jwks_body) + + with gateway_from_environment() as upstream_gateway, wire_server(respond) as jwks: + directory: Final = tmp_path_factory.mktemp("rag_query_deny_by_default") + with owned_proxy( + upstream_gateway, + directory, + {**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url}, + config=_strict_config(directory), + remove_environment=REMOVE_OPENAI_API_BASE, + ) as gateway: + yield StrictGateway(gateway, upstream_gateway, signing_key) + + +StrictCase: TypeAlias = Literal[ + "standalone_key_no_permission", + "team_key_empty_key_grants", + "team_key_empty_team_grants", + "multi_store_one_ungranted", + "rag_alias_no_permission", + "chat_retrieval_config_no_permission", + "jwt_user_without_grant", +] + + +def _strict_denied_request( + strict: StrictGateway, scenario: Scenario, case: StrictCase, model: str, marker: str, store_id: str +) -> tuple[httpx.Response, str]: + models: Final = _json_array(model) + granted: Final = _permission_for_stores(store_id) + empty: Final = _permission_for_stores() + if case == "standalone_key_no_permission": + key = scenario.key(models=models) + return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + "key_vector_store_access_denied" + ) + if case in ("team_key_empty_key_grants", "team_key_empty_team_grants"): + key_grants_store: Final = case == "team_key_empty_team_grants" + team = scenario.team(models=models, object_permission=empty if key_grants_store else granted) + key = scenario.key(team_id=team, models=models, object_permission=granted if key_grants_store else empty) + error_type: Final = "team_vector_store_access_denied" if key_grants_store else "key_vector_store_access_denied" + return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + error_type + ) + if case == "multi_store_one_ungranted": + key = scenario.key(models=models, object_permission=_permission_for_stores(store_id)) + body: Final[JsonObject] = { + **_rag_query_body(model, marker, store_id), + "tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(CONFIG_STORE_ID)}), + } + return strict.gateway.request("POST", "/v1/chat/completions", body, key=key), "key_vector_store_access_denied" + if case == "rag_alias_no_permission": + key = scenario.key(models=models) + return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + "key_vector_store_access_denied" + ) + if case == "chat_retrieval_config_no_permission": + key = scenario.key(models=models) + return ( + strict.gateway.request("POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=key), + "key_vector_store_access_denied", + ) + user: Final = scenario.user(user_role="internal_user") + return ( + strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=strict.jwt(user)), + "user_vector_store_access_denied", + ) + + +@pytest.mark.parametrize( + "case", + ( + "standalone_key_no_permission", + "team_key_empty_key_grants", + "team_key_empty_team_grants", + "multi_store_one_ungranted", + "rag_alias_no_permission", + "chat_retrieval_config_no_permission", + "jwt_user_without_grant", + ), +) +def test_deny_by_default_rejects_ungranted_store_before_upstream_search( + strict_gateway: StrictGateway, case: StrictCase +) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}" + + response, error_type = _strict_denied_request(strict_gateway, scenario, case, model, marker, store_id) + observations: Final = tuple( + observation + for observation in upstream_observations(strict_gateway.upstream) + if marker in str(observation["body"]) + ) + + assert response.status_code == 401, f"{response.text}; scripted_upstream_observations={observations!r}" + assert response.json()["error"]["type"] == error_type, response.text + assert observations == () + + +GrantedCase: TypeAlias = Literal[ + "standalone_key_granted_registered_store", + "team_key_both_grant_unregistered_store", + "jwt_team_member_team_grant_only", + "jwt_user_personal_grant", + "master_key_without_grants", +] + + +def _strict_granted_request( + strict: StrictGateway, scenario: Scenario, case: GrantedCase, model: str, marker: str, store_id: str +) -> httpx.Response: + models: Final = _json_array(model) + granted: Final = _permission_for_stores(store_id) + body: Final = _rag_query_body(model, marker, store_id) + if case in ("standalone_key_granted_registered_store", "team_key_both_grant_unregistered_store"): + team = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None + key = scenario.key(models=models, object_permission=granted, **({} if team is None else {"team_id": team})) + return strict.gateway.request("POST", "/v1/rag/query", body, key=key) + if case == "jwt_team_member_team_grant_only": + team = scenario.team(models=models, object_permission=granted) + member: Final = scenario.member(team) + return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (team,))) + if case == "jwt_user_personal_grant": + user: Final = scenario.user(user_role="internal_user", object_permission=granted) + return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(user)) + return strict.gateway.request("POST", "/v1/rag/query", body) + + +@pytest.mark.parametrize( + "case", + ( + "standalone_key_granted_registered_store", + "team_key_both_grant_unregistered_store", + "jwt_team_member_team_grant_only", + "jwt_user_personal_grant", + "master_key_without_grants", + ), +) +def test_deny_by_default_searches_explicitly_granted_store(strict_gateway: StrictGateway, case: GrantedCase) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = ( + CONFIG_STORE_ID + if case == "standalone_key_granted_registered_store" + else f"vs_unregistered_{uuid.uuid4().hex}" + ) + marker: Final = f"lit6035 granted {case} {uuid.uuid4().hex}" + + response: Final = _strict_granted_request(strict_gateway, scenario, case, model, marker, store_id) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(strict_gateway.upstream, marker, store_id) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + @pytest.mark.parametrize( ("scope", "error_type"), (("key", "key_vector_store_access_denied"), ("team", "team_vector_store_access_denied")), @@ -220,3 +429,31 @@ def test_rag_query_alias_denies_store_when_team_allowlist_excludes(gateway: Gate assert response.status_code == 401, response.text assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text assert _searches_for_marker(gateway, marker) == () + + +def test_explicit_false_flag_keeps_legacy_vector_store_outcomes(tmp_path: Path) -> None: + with gateway_from_environment() as upstream_gateway: + with owned_proxy( + upstream_gateway, + tmp_path, + _openai_environment(upstream_gateway), + config=_strict_config(tmp_path, deny_by_default=False), + remove_environment=REMOVE_OPENAI_API_BASE, + ) as gateway: + with gateway.scenario() as scenario: + model: Final = scenario.model() + allowed_marker: Final = f"flag-false-allowed-{uuid.uuid4().hex}" + denied_marker: Final = f"flag-false-denied-{uuid.uuid4().hex}" + no_permission_key: Final = scenario.key(models=_json_array(model)) + excluding_key: Final = scenario.key( + models=_json_array(model), object_permission=_permission_for_stores("vs_some_other_store") + ) + + allowed: Final = _rag_query(gateway, model, allowed_marker, no_permission_key) + denied: Final = _rag_query(gateway, model, denied_marker, excluding_key) + + assert allowed.status_code == 200, allowed.text + assert len(_searches_for_marker(upstream_gateway, allowed_marker)) == 1 + assert denied.status_code == 401, denied.text + assert denied.json()["error"]["type"] == "key_vector_store_access_denied" + assert _searches_for_marker(upstream_gateway, denied_marker) == () 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 436a773b85f..18c73179dda 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 @@ -1846,6 +1846,8 @@ async def _common_checks_for_rag_query( team_permission: SimpleNamespace | None = None, user_object: LiteLLM_UserTable | None = None, user_permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None = None, + request_body: Mapping[str, object] | None = None, + vector_store_registry: VectorStoreRegistry | None = None, ) -> bool: permissions: Final = { "key-permission": key_permission, @@ -1857,21 +1859,25 @@ async def _common_checks_for_rag_query( side_effect=lambda where: permissions.get(where["object_permission_id"]) ) mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) - request_body = { - "model": "gpt-4o-mini", - "messages": [{"role": "user", "content": "what is in this KB?"}], - "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, - } + body: Final = ( + dict(request_body) + if request_body is not None + else { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + } + ) with ( patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists "litellm.proxy.proxy_server.prisma_client", mock_prisma_client ), patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists - "litellm.vector_store_registry", None + "litellm.vector_store_registry", vector_store_registry ), ): return await common_checks( - request_body=request_body, + request_body=body, team_object=team_object, user_object=user_object, end_user_object=None, @@ -2171,6 +2177,57 @@ async def test_user_owned_standalone_key_cannot_use_owner_grants_under_deny_by_d ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("deny_by_default", "key_vector_stores", "denied_by"), + [ + (True, ["KBSTOREA", "KBSTOREB"], None), + (True, ["KBSTOREA"], _KEY_DENIED), + (True, ["KBSTOREB"], _KEY_DENIED), + (False, ["KBSTOREA"], _KEY_DENIED), + ], + ids=["enabled-grants-both", "enabled-missing-second", "enabled-missing-first", "disabled-missing-second"], +) +async def test_every_requested_vector_store_needs_a_grant( + deny_by_default: bool, key_vector_stores: list[str], denied_by: ProxyErrorTypes | None +): + 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"), + SimpleNamespace(vector_stores=key_vector_stores), + request_body={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in these KBs?"}], + "tools": [{"type": "file_search", "vector_store_ids": ["KBSTOREB"]}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + }, + vector_store_registry=VectorStoreRegistry(), + ), + denied_by, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "valid_token", + [_virtual_key(), UserAPIKeyAuth(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER)], + ids=["standalone-key", "keyless-user"], +) +async def test_request_without_vector_stores_is_unaffected_by_deny_by_default(valid_token: UserAPIKeyAuth): + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + valid_token, + None, + user_object=_user_row(None), + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]}, + vector_store_registry=VectorStoreRegistry(), + ), + None, + ) + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router From e82721fd72411f5e55b24eb9fceb6b1b2f170cb0 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 20:36:36 +0000 Subject: [PATCH 05/14] chore(ui): regenerate schema.d.ts for vector_store_deny_by_default Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index a376f8d7910..3f47d7b5ac7 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29897,6 +29897,12 @@ export interface components { * @description Master switch for the SSRF guard applied to user-supplied URLs (image_url, file_url, MCP/OpenAPI spec URLs, etc). Defaults to True. Set to False to disable DNS/IP validation entirely (not recommended). */ user_url_validation?: boolean | null; + /** + * Vector Store Deny By Default + * @description When True, a vector store must be explicitly listed in object_permission.vector_stores: a virtual key needs its own grant plus its team's, a keyless team member needs the team's, and a user with neither needs their own. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys are not yet covered + * @default false + */ + vector_store_deny_by_default: boolean; }; /** ConfigList */ ConfigList: { From bda87054fcd52debbca01c8c207bf4a2ca87bcb2 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 21:57:40 +0000 Subject: [PATCH 06/14] fix(auth): cover path and file_search vector store ids under vector_store_deny_by_default Strict mode now reads vector_store_ids and tools[].vector_store_ids from the request body without needing a vector store registry, so /v1/vector_stores/{id}/search and Responses file_search are checked. User grants load through the object permission cache, and the proxy admin user rebuild keeps object_permission_id Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 52 ++++++-- litellm/proxy/auth/user_api_key_auth.py | 3 + .../test_rag_query_vector_store_allowlist.py | 116 +++++++++++++++--- ...st_auth_checks_object_access_and_lookup.py | 84 +++++++++++-- .../test_user_api_key_auth_request_flow.py | 91 ++++++++++++-- .../proxy/proxy_server/test_proxy_config.py | 2 +- 6 files changed, 302 insertions(+), 46 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 9a4f72441cb..475ad5bbba9 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -15,12 +15,13 @@ import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from functools import partial +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias from fastapi import HTTPException, Request, status from pydantic import BaseModel, TypeAdapter, ValidationError -from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack +from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, TypeIs, Unpack import litellm from litellm._logging import verbose_proxy_logger @@ -6635,6 +6636,32 @@ 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_object_list(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing + return isinstance(value, list) + + +def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: request JSON keys are str + return isinstance(value, dict) + + +def _object_items(value: object) -> tuple[object, ...]: + return tuple(value) if _is_object_list(value) else () + + +def _tool_vector_store_ids(tool: object) -> tuple[object, ...]: + return _object_items(tool.get("vector_store_ids")) if _is_object_mapping(tool) else () + + +def _get_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]: + candidate_ids: Final = ( + *_object_items(request_body.get("vector_store_ids")), + *chain.from_iterable(_tool_vector_store_ids(tool) for tool in _object_items(request_body.get("tools"))), + ) + return tuple( + vector_store_id for vector_store_id in candidate_ids if isinstance(vector_store_id, str) and vector_store_id + ) + + def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool: return ( valid_token is not None @@ -6675,7 +6702,7 @@ async def vector_store_access_check( Raises ProxyException if the object (key, team, org) cannot access the specific vector store. """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache ######################################################### # Get the vector store the user is trying to access @@ -6685,12 +6712,17 @@ async def vector_store_access_check( return True registry_ids: Final = ( - litellm.vector_store_registry.get_vector_store_ids_to_run( - non_default_params=request_body, tools=request_body.get("tools", None) + _get_requested_vector_store_ids(_typed_request_body(request_body)) + if deny_by_default + else ( + litellm.vector_store_registry.get_vector_store_ids_to_run( + non_default_params=request_body, tools=request_body.get("tools", None) + ) + if litellm.vector_store_registry is not None + else None ) - if litellm.vector_store_registry is not None - else None - ) or () + or () + ) rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body)) rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else () vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids))) @@ -6742,8 +6774,10 @@ async def vector_store_access_check( if strict_identity and not strict_key and not has_team: user_object_permission: Final = ( - await _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique( - where={"object_permission_id": user_object.object_permission_id}, + await get_object_permission( + object_permission_id=user_object.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, ) if user_object is not None and user_object.object_permission_id is not None else None diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e82f3eed7cc..a8aec448b5b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3054,6 +3054,9 @@ async def _run_centralized_common_checks( user_id=user_api_key_auth_obj.user_id or litellm_proxy_admin_name, user_role=LitellmUserRoles.PROXY_ADMIN, spend=user_object.spend if user_object is not None else 0.0, + object_permission_id=( + user_object.object_permission_id if isinstance(user_object, LiteLLM_UserTable) else None + ), ) if project_object is not None: diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index 4f1270a33aa..a5807dba3b7 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -182,34 +182,44 @@ def _strict_denied_request( granted: Final = _permission_for_stores(store_id) empty: Final = _permission_for_stores() if case == "standalone_key_no_permission": - key = scenario.key(models=models) - return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + standalone_key: Final = scenario.key(models=models) + return strict.gateway.request( + "POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=standalone_key + ), ( "key_vector_store_access_denied" ) if case in ("team_key_empty_key_grants", "team_key_empty_team_grants"): key_grants_store: Final = case == "team_key_empty_team_grants" - team = scenario.team(models=models, object_permission=empty if key_grants_store else granted) - key = scenario.key(team_id=team, models=models, object_permission=granted if key_grants_store else empty) + team: Final = scenario.team(models=models, object_permission=empty if key_grants_store else granted) + team_key: Final = scenario.key( + team_id=team, models=models, object_permission=granted if key_grants_store else empty + ) error_type: Final = "team_vector_store_access_denied" if key_grants_store else "key_vector_store_access_denied" - return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + return strict.gateway.request( + "POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=team_key + ), ( error_type ) if case == "multi_store_one_ungranted": - key = scenario.key(models=models, object_permission=_permission_for_stores(store_id)) + partial_key: Final = scenario.key(models=models, object_permission=_permission_for_stores(store_id)) body: Final[JsonObject] = { **_rag_query_body(model, marker, store_id), "tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(CONFIG_STORE_ID)}), } - return strict.gateway.request("POST", "/v1/chat/completions", body, key=key), "key_vector_store_access_denied" + return strict.gateway.request( + "POST", "/v1/chat/completions", body, key=partial_key + ), "key_vector_store_access_denied" if case == "rag_alias_no_permission": - key = scenario.key(models=models) - return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=key), ( + alias_key: Final = scenario.key(models=models) + return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=alias_key), ( "key_vector_store_access_denied" ) if case == "chat_retrieval_config_no_permission": - key = scenario.key(models=models) + chat_key: Final = scenario.key(models=models) return ( - strict.gateway.request("POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=key), + strict.gateway.request( + "POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=chat_key + ), "key_vector_store_access_denied", ) user: Final = scenario.user(user_role="internal_user") @@ -267,13 +277,15 @@ def _strict_granted_request( granted: Final = _permission_for_stores(store_id) body: Final = _rag_query_body(model, marker, store_id) if case in ("standalone_key_granted_registered_store", "team_key_both_grant_unregistered_store"): - team = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None - key = scenario.key(models=models, object_permission=granted, **({} if team is None else {"team_id": team})) - return strict.gateway.request("POST", "/v1/rag/query", body, key=key) + key_team: Final = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None + granted_key: Final = scenario.key( + models=models, object_permission=granted, **({} if key_team is None else {"team_id": key_team}) + ) + return strict.gateway.request("POST", "/v1/rag/query", body, key=granted_key) if case == "jwt_team_member_team_grant_only": - team = scenario.team(models=models, object_permission=granted) - member: Final = scenario.member(team) - return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (team,))) + member_team: Final = scenario.team(models=models, object_permission=granted) + member: Final = scenario.member(member_team) + return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (member_team,))) if case == "jwt_user_personal_grant": user: Final = scenario.user(user_role="internal_user", object_permission=granted) return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(user)) @@ -457,3 +469,73 @@ def test_explicit_false_flag_keeps_legacy_vector_store_outcomes(tmp_path: Path) assert denied.status_code == 401, denied.text assert denied.json()["error"]["type"] == "key_vector_store_access_denied" assert _searches_for_marker(upstream_gateway, denied_marker) == () + + +def _strict_no_registry_config(directory: Path) -> Path: + config: Final = object_value(yaml.safe_load(_no_registry_config(directory).read_text())) + general_settings: Final = object_value(config["general_settings"]) + strict: Final[JsonObject] = { + **config, + "general_settings": {**general_settings, "vector_store_deny_by_default": True}, + } + path: Final = directory / "proxy_vector_store_deny_by_default_no_registry.yaml" + path.write_text(yaml.safe_dump(strict)) + return path + + +def test_deny_by_default_without_registry_checks_search_route_and_file_search_tools(tmp_path: Path) -> None: + with gateway_from_environment() as upstream_gateway: + with owned_proxy( + upstream_gateway, + tmp_path, + _openai_environment(upstream_gateway), + config=_strict_no_registry_config(tmp_path), + remove_environment=REMOVE_OPENAI_API_BASE, + ) as gateway: + with gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + search_marker: Final = f"lit6035 no registry search {uuid.uuid4().hex}" + responses_marker: Final = f"lit6035 no registry responses {uuid.uuid4().hex}" + granted_marker: Final = f"lit6035 no registry granted search {uuid.uuid4().hex}" + ungranted_key: Final = scenario.key(models=_json_array(model)) + granted_key: Final = scenario.key( + models=_json_array(model), object_permission=_permission_for_stores(store_id) + ) + + search_denied: Final = gateway.request( + "POST", f"/v1/vector_stores/{store_id}/search", {"query": search_marker}, key=ungranted_key + ) + responses_denied: Final = gateway.request( + "POST", + "/v1/responses", + { + "model": model, + "input": responses_marker, + "tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(store_id)}), + }, + key=ungranted_key, + ) + search_granted: Final = gateway.request( + "POST", f"/v1/vector_stores/{store_id}/search", {"query": granted_marker}, key=granted_key + ) + + observations: Final = upstream_observations(upstream_gateway) + denied_observations: Final = tuple( + observation + for observation in observations + if search_marker in str(observation["body"]) or responses_marker in str(observation["body"]) + ) + granted_searches: Final = tuple( + observation + for observation in observations + if observation["path"] == f"/vector_stores/{store_id}/search" + and granted_marker in str(observation["body"]) + ) + assert search_denied.status_code == 401, search_denied.text + assert search_denied.json()["error"]["type"] == "key_vector_store_access_denied", search_denied.text + assert responses_denied.status_code == 401, responses_denied.text + assert responses_denied.json()["error"]["type"] == "key_vector_store_access_denied", responses_denied.text + assert denied_observations == () + assert search_granted.status_code == 200, search_granted.text + assert len(granted_searches) == 1, granted_searches 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 18c73179dda..03164f0e4d7 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 @@ -92,6 +92,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, end_user_cache_key, end_user_restricted_registry_cache_key, + object_permission_cache_key, tag_cache_key, tag_registry_cache_key, ) @@ -1827,7 +1828,7 @@ def _virtual_key( user_role: LitellmUserRoles | None = None, team_id: str | None = None, ) -> UserAPIKeyAuth: - key = UserAPIKeyAuth( + key: Final = UserAPIKeyAuth( api_key=api_key, user_id="key-owner", user_role=user_role, @@ -1847,14 +1848,21 @@ async def _common_checks_for_rag_query( user_object: LiteLLM_UserTable | None = None, user_permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None = None, request_body: Mapping[str, object] | None = None, + user_api_key_cache: UserApiKeyCache | None = None, vector_store_registry: VectorStoreRegistry | None = None, ) -> bool: permissions: Final = { "key-permission": key_permission, "team-permission": team_permission, - "user-permission": user_permission, + "user-permission": ( + LiteLLM_ObjectPermissionTable( + object_permission_id="user-permission", vector_stores=user_permission.vector_stores + ) + if isinstance(user_permission, SimpleNamespace) + else user_permission + ), } - mock_prisma_client = MagicMock() + mock_prisma_client: Final = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( side_effect=lambda where: permissions.get(where["object_permission_id"]) ) @@ -1875,6 +1883,10 @@ async def _common_checks_for_rag_query( patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists "litellm.vector_store_registry", vector_store_registry ), + patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists + "litellm.proxy.proxy_server.user_api_key_cache", + UserApiKeyCache() if user_api_key_cache is None else user_api_key_cache, + ), ): return await common_checks( request_body=body, @@ -1958,7 +1970,7 @@ async def test_standalone_key_vector_store_permission_record_under_deny_by_defau @pytest.mark.asyncio @pytest.mark.parametrize("deny_by_default", [False, True], ids=["flag-false", "flag-true"]) async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_default: bool): - master_key = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN) + master_key: Final = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN) await _assert_rag_query_outcome( _common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), None @@ -1978,10 +1990,10 @@ async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_defa async def test_proxy_admin_virtual_key_keeps_flag_off_vector_store_behavior( general_settings: Mapping[str, object], vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None ): - admin_key = _virtual_key( + admin_key: Final = _virtual_key( object_permission_id=None if vector_stores is None else "key-permission", user_role=LitellmUserRoles.PROXY_ADMIN ) - key_permission = None if vector_stores is None else SimpleNamespace(vector_stores=vector_stores) + 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) @@ -2033,7 +2045,7 @@ 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 = 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 @@ -2136,7 +2148,7 @@ async def test_keyless_team_member_needs_only_team_grant_under_deny_by_default( async def test_session_token_is_not_a_virtual_key_under_deny_by_default( team_id: str | None, denied_by: ProxyErrorTypes | None ): - session = UserAPIKeyAuth( + session: Final = UserAPIKeyAuth( api_key="hashed-session-token", user_id="user-1", team_id=team_id, @@ -2228,6 +2240,62 @@ async def test_request_without_vector_stores_is_unaffected_by_deny_by_default(va ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "request_body", + [ + {"query": "what is in this KB?", "vector_store_id": "KBSTOREA", "vector_store_ids": ["KBSTOREA"]}, + { + "model": "gpt-4o-mini", + "input": "what is in this KB?", + "tools": [{"type": "file_search", "vector_store_ids": ["KBSTOREA"]}], + }, + ], + ids=["vector-store-search-route", "responses-file-search"], +) +@pytest.mark.parametrize( + ("deny_by_default", "key_vector_stores", "denied_by"), + [(True, None, _KEY_DENIED), (True, ["KBSTOREA"], None), (False, None, None)], + ids=["enabled-no-grant", "enabled-grant", "disabled-no-grant"], +) +async def test_deny_by_default_reads_requested_vector_stores_without_a_registry( + request_body: Mapping[str, object], + deny_by_default: bool, + key_vector_stores: list[str] | None, + denied_by: ProxyErrorTypes | None, +): + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": deny_by_default}, + _virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission"), + None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores), + request_body=request_body, + ), + denied_by, + ) + + +@pytest.mark.asyncio +async def test_keyless_user_grant_is_read_through_the_object_permission_cache(): + cache: Final = UserApiKeyCache(default_in_memory_ttl=60) + await cache.async_set_cache( + key=object_permission_cache_key("user-permission"), + value=LiteLLM_ObjectPermissionTable(object_permission_id="user-permission", vector_stores=["KBSTOREA"]), + model_type=LiteLLM_ObjectPermissionTable, + ) + + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + UserAPIKeyAuth(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER), + None, + user_object=_user_row(["KBSTOREA"]), + user_api_key_cache=cache, + ), + None, + ) + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index eb17ff0823c..76f7fb2aae6 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -12,6 +12,7 @@ from functools import partial from pathlib import Path from textwrap import dedent from types import SimpleNamespace +from typing import Final from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -21,6 +22,7 @@ from fastapi import HTTPException, status import litellm import litellm.proxy.proxy_server from litellm.caching.dual_cache import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy._types import ( LiteLLMRoutes, LiteLLM_JWTAuth, @@ -5318,19 +5320,19 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( from fastapi import Request from starlette.datastructures import URL - token = UserAPIKeyAuth( + token: Final = UserAPIKeyAuth( api_key="sk-team-key", team_id="team-1", team_models=["gpt-4o-mini"], object_permission_id="key-permission" ) token.via_virtual_key = True - request = Request(scope={"type": "http"}) + request: Final = Request(scope={"type": "http"}) request._url = URL(url="/v1/rag/query") - database = MagicMock() + database: Final = MagicMock() database.db.litellm_objectpermissiontable.find_unique = AsyncMock( side_effect=lambda where: SimpleNamespace(vector_stores=["KBSTOREA"]) if where["object_permission_id"] == "key-permission" else None ) - attrs = { + attrs: Final = { **_proxy_attrs_for_centralized_checks(), "prisma_client": database, "general_settings": {"vector_store_deny_by_default": deny_by_default}, @@ -5342,7 +5344,7 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( "litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error) ) - checks = _run_centralized_common_checks( + checks: Final = _run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5380,25 +5382,25 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te from fastapi import Request from starlette.datastructures import URL - token = UserAPIKeyAuth( + token: Final = UserAPIKeyAuth( user_id="user-1", team_id="team-1", team_models=["gpt-4o-mini"], user_role=LitellmUserRoles.INTERNAL_USER, ) - request = Request(scope={"type": "http"}) + request: Final = Request(scope={"type": "http"}) request._url = URL(url="/v1/rag/query") - grants = { + grants: Final = { "team-1-grants-a": ["KBSTOREA"], "team-1-grants-none": [], "team-2-grants-a": ["KBSTOREA"], "user-grants-a": ["KBSTOREA"], } - database = MagicMock() + database: Final = MagicMock() database.db.litellm_objectpermissiontable.find_unique = AsyncMock( side_effect=lambda where: SimpleNamespace(vector_stores=grants[where["object_permission_id"]]) ) - attrs = { + attrs: Final = { **_proxy_attrs_for_centralized_checks(), "prisma_client": database, "general_settings": {"vector_store_deny_by_default": True}, @@ -5423,7 +5425,7 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te ), ) - checks = _run_centralized_common_checks( + checks: Final = _run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5445,6 +5447,73 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_permission_id", "denied"), + [("admin-grants-a", False), (None, True)], + ids=["admin-personal-grant", "admin-without-grant"], +) +async def test_keyless_proxy_admin_keeps_personal_vector_store_grants_under_deny_by_default( + monkeypatch: pytest.MonkeyPatch, user_permission_id: str | None, denied: bool +): + from fastapi import Request + from starlette.datastructures import URL + + token: Final = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN) + request: Final = Request(scope={"type": "http"}) + request._url = URL(url="/v1/rag/query") + database: Final = MagicMock() + database.db.litellm_objectpermissiontable.find_unique = AsyncMock( + side_effect=lambda where: SimpleNamespace( + dict=lambda: {"object_permission_id": where["object_permission_id"], "vector_stores": ["KBSTOREA"]}, + vector_stores=["KBSTOREA"], + ) + ) + attrs: Final = { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "general_settings": {"vector_store_deny_by_default": True}, + "user_api_key_cache": UserApiKeyCache(), + "proxy_logging_obj": MagicMock( + service_logging_obj=MagicMock( + async_service_success_hook=AsyncMock(), async_service_failure_hook=AsyncMock() + ) + ), + } + for name, value in attrs.items(): + monkeypatch.setattr(litellm.proxy.proxy_server, name, value) + monkeypatch.setattr(litellm, "vector_store_registry", None) + monkeypatch.setattr( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + AsyncMock( + return_value=LiteLLM_UserTable( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN, object_permission_id=user_permission_id + ) + ), + ) + + checks: Final = _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"}, + }, + route="/v1/rag/query", + ) + if not denied: + await checks + return + with pytest.raises(ProxyException) as exc_info: + await checks + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( + ProxyErrorTypes.user_vector_store_access_denied, + "vector_store", + "401", + ) + + @pytest.mark.asyncio async def test_centralized_common_checks_propagates_end_user_budget_error(): """Regression: ``get_end_user_object`` raises ``litellm.BudgetExceededError`` diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 0c7a4a1608b..804a372a1c9 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -2279,7 +2279,7 @@ async def test_load_config_logs_disabled_budget_reservation_once(tmp_path, monke async def test_load_config_yaml_vector_store_deny_by_default_is_boolean( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str, expected: bool ): - config_file = tmp_path / "vector_store.yaml" + config_file: Final = tmp_path / "vector_store.yaml" config_file.write_text( f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n" ) From 478e845f9ced36e0a0f7b83fef38f9a538297c23 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 22:20:40 +0000 Subject: [PATCH 07/14] fix(proxy): broadcast user entitlement cache eviction to every worker Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../internal_user_endpoints.py | 6 +- .../test_rag_query_vector_store_allowlist.py | 67 ++++++++++++++++--- .../test_internal_user_endpoints.py | 28 ++++++++ 3 files changed, 85 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 04b8ec56ae2..cbdd4051297 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1451,11 +1451,7 @@ async def _invalidate_cached_user_entitlement(user_id: str | None, object_permis *(object_permission_cache_key(permission_id) for permission_id in dict.fromkeys(object_permission_ids)), *((user_object_permission_id_cache_key(user_id), user_id) if user_id is not None else ()), ) - for key in keys: - try: - await user_api_key_cache.async_delete_cache(key=key) - except Exception as e: # noqa: BLE001 # a cache we cannot clear still expires; never fail the write - verbose_proxy_logger.warning("Failed to invalidate cached entitlement key %r: %s", key, e) + await evict_and_broadcast(cache_keys=keys, user_api_key_cache=user_api_key_cache) async def _update_single_user_helper( diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index a5807dba3b7..bd295f9e822 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -1,6 +1,7 @@ from __future__ import annotations import json +import os import time import uuid from collections.abc import Iterator, Mapping @@ -13,16 +14,18 @@ import jwt import pytest import yaml from cryptography.hazmat.primitives.asymmetric import rsa -from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server from integration.authorization._guardrail_opt_out import upstream_observations from pydantic import JsonValue +from redis import Redis CONFIG_STORE_ID: Final = "vs_integration_config_store" PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",) JWT_KEY_ID: Final = "integration-vector-store-jwt-key" +AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation" JsonObject: TypeAlias = dict[str, JsonValue] @@ -106,10 +109,19 @@ def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[t class StrictGateway: - def __init__(self, gateway: Gateway, upstream: Gateway, signing_key: rsa.RSAPrivateKey) -> None: + def __init__( + self, + gateway: Gateway, + upstream: Gateway, + signing_key: rsa.RSAPrivateKey, + config: Path, + environment: Mapping[str, str], + ) -> None: self.gateway: Final = gateway self.upstream: Final = upstream self._signing_key: Final = signing_key + self.config: Final = config + self.environment: Final = environment def jwt(self, subject: str, groups: tuple[str, ...] = ()) -> str: claims: Final[JsonObject] = { @@ -154,14 +166,16 @@ def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictG with gateway_from_environment() as upstream_gateway, wire_server(respond) as jwks: directory: Final = tmp_path_factory.mktemp("rag_query_deny_by_default") + config: Final = _strict_config(directory) + environment: Final = MappingProxyType({**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url}) with owned_proxy( upstream_gateway, directory, - {**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url}, - config=_strict_config(directory), + environment, + config=config, remove_environment=REMOVE_OPENAI_API_BASE, ) as gateway: - yield StrictGateway(gateway, upstream_gateway, signing_key) + yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment) StrictCase: TypeAlias = Literal[ @@ -185,9 +199,7 @@ def _strict_denied_request( standalone_key: Final = scenario.key(models=models) return strict.gateway.request( "POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=standalone_key - ), ( - "key_vector_store_access_denied" - ) + ), ("key_vector_store_access_denied") if case in ("team_key_empty_key_grants", "team_key_empty_team_grants"): key_grants_store: Final = case == "team_key_empty_team_grants" team: Final = scenario.team(models=models, object_permission=empty if key_grants_store else granted) @@ -197,9 +209,7 @@ def _strict_denied_request( error_type: Final = "team_vector_store_access_denied" if key_grants_store else "key_vector_store_access_denied" return strict.gateway.request( "POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=team_key - ), ( - error_type - ) + ), (error_type) if case == "multi_store_one_ungranted": partial_key: Final = scenario.key(models=models, object_permission=_permission_for_stores(store_id)) body: Final[JsonObject] = { @@ -539,3 +549,38 @@ def test_deny_by_default_without_registry_checks_search_route_and_file_search_to assert denied_observations == () assert search_granted.status_code == 200, search_granted.text assert len(granted_searches) == 1, granted_searches + + +def _auth_cache_subscribers() -> int: + with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache: + return int(cache.pubsub_numsub(AUTH_CACHE_INVALIDATION_CHANNEL)[0][1]) + + +def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: StrictGateway, tmp_path: Path) -> None: + subscribers_before_peer: Final = _auth_cache_subscribers() + with ( + owned_proxy( + strict_gateway.upstream, + tmp_path, + strict_gateway.environment, + config=strict_gateway.config, + remove_environment=REMOVE_OPENAI_API_BASE, + ) as peer, + strict_gateway.gateway.scenario() as scenario, + ): + eventually(_auth_cache_subscribers, lambda count: count > subscribers_before_peer) + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + user: Final = scenario.user(user_role="internal_user", object_permission=_permission_for_stores(store_id)) + token: Final = strict_gateway.jwt(user) + + def peer_status() -> int: + marker: Final = f"lit6035 revoked user grant {uuid.uuid4().hex}" + return peer.request( + "POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=token + ).status_code + + assert eventually(peer_status, lambda status: status == 200, seconds=30, return_last_on_timeout=True) == 200 + strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()}) + + assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401 diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 799fe59147c..a253dcfb5d8 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -4003,6 +4003,34 @@ async def test_user_update_invalidates_the_cached_entitlement(mocker): } +@pytest.mark.asyncio +async def test_user_update_broadcasts_the_entitlement_invalidation_to_other_workers(mocker: MockerFixture): + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + _object_permission_mocks(mocker) + cache: Final = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: substitute the cache dependency + broadcast: Final = mocker.patch( # test-quality-ok: observe the Redis publication boundary + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", + new_callable=mocker.AsyncMock, + ) + + await _update_single_user_helper( + user_request=UpdateUserRequest(user_id="target-user", object_permission={"vector_stores": []}), + user_api_key_dict=UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + broadcast_keys: Final = {call.kwargs["cache_key"] for call in broadcast.await_args_list} + assert broadcast_keys == { + "object_permission_id:perm-new", + "user_object_permission_id:target-user", + "target-user", + } + + @pytest.mark.asyncio async def test_admin_can_clear_a_users_mcp_entitlement(mocker): """An explicit empty object_permission means "no object permission", so it must unlink. From ad1e28767e64016ac9a748c6fe3a964eacf3b8d7 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 22:53:54 +0000 Subject: [PATCH 08/14] refactor(auth): reuse VectorStoreRegistry id extraction for vector_store_deny_by_default Strict mode now uses get_vector_store_ids_to_run on an empty registry when none is loaded, instead of a parallel set of request-shape helpers, and vector_store_access_check documents the policy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 76 ++++++++++++++----------------- 1 file changed, 35 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 475ad5bbba9..30f8349b79d 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -15,13 +15,12 @@ import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from functools import partial -from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias from fastapi import HTTPException, Request, status from pydantic import BaseModel, TypeAdapter, ValidationError -from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, TypeIs, Unpack +from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack import litellm from litellm._logging import verbose_proxy_logger @@ -147,6 +146,7 @@ from litellm.router import Router from litellm.types.proxy.auth.auth_checks import UserNotFoundError from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget from litellm.utils import get_utc_datetime +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry from .auth_checks_organization import ( add_team_org_context_to_request_body, @@ -6636,32 +6636,6 @@ 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_object_list(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing - return isinstance(value, list) - - -def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: request JSON keys are str - return isinstance(value, dict) - - -def _object_items(value: object) -> tuple[object, ...]: - return tuple(value) if _is_object_list(value) else () - - -def _tool_vector_store_ids(tool: object) -> tuple[object, ...]: - return _object_items(tool.get("vector_store_ids")) if _is_object_mapping(tool) else () - - -def _get_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]: - candidate_ids: Final = ( - *_object_items(request_body.get("vector_store_ids")), - *chain.from_iterable(_tool_vector_store_ids(tool) for tool in _object_items(request_body.get("tools"))), - ) - return tuple( - vector_store_id for vector_store_id in candidate_ids if isinstance(vector_store_id, str) and vector_store_id - ) - - def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool: return ( valid_token is not None @@ -6698,9 +6672,29 @@ async def vector_store_access_check( deny_by_default: bool = False, ): """ - Checks if the object (key, team, org) has access to the vector store. + Checks whether the caller may use every vector store the request names. - Raises ProxyException if the object (key, team, org) cannot access the specific vector store. + Requested stores come from `vector_store_ids`, `tools[].vector_store_ids` and the RAG + `retrieval_config.vector_store_id`. Grants come from each identity's + `object_permission.vector_stores`. + + With `deny_by_default=False` (legacy), a key or team only restricts access when its list is + nonempty, and stores are only read from the request when a vector store registry is loaded. + + With `deny_by_default=True` (`general_settings.vector_store_deny_by_default`), stores are read + even without a registry, and a missing record, `null` or `[]` grants nothing: + + - virtual key: the key must grant every store, and so must its team when it has one, even if + the team failed to load + - keyless team member (JWT, `lite login` session token): only the resolved team is checked + - keyless user with no team: the user's own grant is checked + - master key and dashboard sessions: legacy behavior + + The user's personal grant is only consulted in the keyless no-team case, so it can neither + rescue nor restrict a key or team request. + + Raises ProxyException (401, `{key,team,user}_vector_store_access_denied`) on the first identity + that does not grant a requested store. """ from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -6711,18 +6705,18 @@ async def vector_store_access_check( verbose_proxy_logger.debug("Prisma client not found, skipping vector store access check") return True - registry_ids: Final = ( - _get_requested_vector_store_ids(_typed_request_body(request_body)) - if deny_by_default - else ( - litellm.vector_store_registry.get_vector_store_ids_to_run( - non_default_params=request_body, tools=request_body.get("tools", None) - ) - if litellm.vector_store_registry is not None - else None - ) - or () + vector_store_registry: Final = ( + VectorStoreRegistry() + if litellm.vector_store_registry is None and deny_by_default + else litellm.vector_store_registry ) + registry_ids: Final = ( + vector_store_registry.get_vector_store_ids_to_run( + non_default_params=request_body, tools=request_body.get("tools", None) + ) + if vector_store_registry is not None + else None + ) or () rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body)) rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else () vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids))) From 3ccaf280472f96ef8f2bfd5e3ccb9abf49e94462 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 23:45:15 +0000 Subject: [PATCH 09/14] test(integration): add strict vector store audit cells for routes, SDKs, workers and concurrency Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_rag_query_vector_store_allowlist.py | 356 +++++++++++++++++- 1 file changed, 354 insertions(+), 2 deletions(-) diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index bd295f9e822..3ee320994d4 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -1,16 +1,20 @@ from __future__ import annotations +import asyncio import json import os import time import uuid from collections.abc import Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from types import MappingProxyType from typing import Final, Literal, TypeAlias +import anthropic import httpx import jwt +import openai import pytest import yaml from cryptography.hazmat.primitives.asymmetric import rsa @@ -21,6 +25,9 @@ from integration.authorization._guardrail_opt_out import upstream_observations from pydantic import JsonValue from redis import Redis +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + CONFIG_STORE_ID: Final = "vs_integration_config_store" PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",) @@ -66,6 +73,20 @@ def _rag_query( return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key) +def _served_model(gateway: Gateway, scenario: Scenario) -> str: + model: Final = scenario.model() + statuses: Final = eventually( + lambda: frozenset( + _rag_query(gateway, model, f"lit6035 warm {uuid.uuid4().hex}", gateway.key).status_code for _ in range(8) + ), + lambda seen: seen == frozenset({200}), + seconds=30, + return_last_on_timeout=True, + ) + assert statuses == frozenset({200}), statuses + return model + + def _searches_for_marker( gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID ) -> tuple[Mapping[str, JsonValue], ...]: @@ -174,6 +195,7 @@ def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictG environment, config=config, remove_environment=REMOVE_OPENAI_API_BASE, + workers=2, ) as gateway: yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment) @@ -255,7 +277,7 @@ def test_deny_by_default_rejects_ungranted_store_before_upstream_search( strict_gateway: StrictGateway, case: StrictCase ) -> None: with strict_gateway.gateway.scenario() as scenario: - model: Final = scenario.model() + model: Final = _served_model(strict_gateway.gateway, scenario) store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}" @@ -314,7 +336,7 @@ def _strict_granted_request( ) def test_deny_by_default_searches_explicitly_granted_store(strict_gateway: StrictGateway, case: GrantedCase) -> None: with strict_gateway.gateway.scenario() as scenario: - model: Final = scenario.model() + model: Final = _served_model(strict_gateway.gateway, scenario) store_id: Final = ( CONFIG_STORE_ID if case == "standalone_key_granted_registered_store" @@ -584,3 +606,333 @@ def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: Stric strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()}) assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401 + + +def _store_searches(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + observation + for observation in upstream_observations(gateway) + if str(observation["path"]).startswith("/vector_stores/") and marker in str(observation["body"]) + ) + + +def _marker_observations(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(observation for observation in upstream_observations(gateway) if marker in str(observation["body"])) + + +def _messages_body(model: str, marker: str, store_id: str) -> JsonObject: + body: Final[JsonObject] = { + "model": model, + "max_tokens": 16, + "messages": _json_array({"role": "user", "content": marker}), + "vector_store_ids": _json_array(store_id), + } + return body + + +def _file_search_tools(store_id: JsonValue) -> JsonValue: + return _json_array({"type": "file_search", "vector_store_ids": store_id}) + + +RequestShape: TypeAlias = Literal["messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids"] + + +def _shape_request(shape: RequestShape, model: str, marker: str, store_id: str) -> tuple[str, JsonObject]: + messages: Final = _json_array({"role": "user", "content": marker}) + if shape == "messages_vector_store_ids": + return "/v1/messages", _messages_body(model, marker, store_id) + if shape == "chat_stream_file_search": + stream_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "stream": True, + "tools": _file_search_tools(_json_array(store_id)), + } + return "/v1/chat/completions", stream_body + chat_body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": _json_array(store_id)} + return "/v1/chat/completions", chat_body + + +@pytest.mark.parametrize("shape", ("messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids")) +def test_deny_by_default_covers_messages_streaming_and_top_level_store_ids( + strict_gateway: StrictGateway, shape: RequestShape +) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 shape {shape} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + path, body = _shape_request(shape, model, marker, store_id) + + response: Final = strict_gateway.gateway.request("POST", path, body, key=key) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +SdkClient: TypeAlias = Literal["openai_sync", "openai_async", "anthropic_sync"] + + +def _sdk_denial(gateway: Gateway, client: SdkClient, key: str, model: str, marker: str, store_id: str) -> int: + base_url: Final = str(gateway.client.base_url) + if client == "anthropic_sync": + with pytest.raises(anthropic.AuthenticationError) as anthropic_denied: + anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_body={"vector_store_ids": _json_array(store_id)}, + ) + return anthropic_denied.value.status_code + tools: Final = [{"type": "file_search", "vector_store_ids": [store_id]}] + if client == "openai_sync": + with pytest.raises(openai.AuthenticationError) as sync_denied: + openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0).responses.create( + model=model, input=marker, tools=tools + ) + return sync_denied.value.status_code + + async def create() -> None: + async with openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as async_client: + await async_client.responses.create(model=model, input=marker, tools=tools) + + with pytest.raises(openai.AuthenticationError) as async_denied: + asyncio.run(create()) + return async_denied.value.status_code + + +@pytest.mark.parametrize("client", ("openai_sync", "openai_async", "anthropic_sync")) +def test_deny_by_default_rejects_sdk_clients_without_grant(strict_gateway: StrictGateway, client: SdkClient) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 sdk {client} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + + status: Final = _sdk_denial(strict_gateway.gateway, client, key, model, marker, store_id) + + assert status == 401 + assert _marker_observations(strict_gateway.upstream, marker) == () + + +MalformedShape: TypeAlias = Literal[ + "ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id" +] + + +def _malformed_body(shape: MalformedShape, model: str, marker: str) -> JsonObject: + messages: Final = _json_array({"role": "user", "content": marker}) + if shape == "tool_ids_string": + tool_ids_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "tools": _file_search_tools(f"vs_{uuid.uuid4().hex}"), + } + return tool_ids_body + if shape == "tool_is_string": + string_tool_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "tools": _json_array("file_search"), + } + return string_tool_body + ids: Final[Mapping[MalformedShape, JsonValue]] = MappingProxyType( + { + "ids_string": f"vs_{uuid.uuid4().hex}", + "ids_int": _json_array(123), + "ids_empty_string": _json_array(""), + "oversized_id": _json_array("vs_" + "x" * 5000), + } + ) + body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": ids[shape]} + return body + + +@pytest.mark.parametrize( + "shape", ("ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id") +) +def test_deny_by_default_malformed_store_ids_never_search_a_store( + strict_gateway: StrictGateway, shape: MalformedShape +) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + marker: Final = f"lit6035 malformed {shape} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + + response: Final = strict_gateway.gateway.request( + "POST", "/v1/chat/completions", _malformed_body(shape, model, marker), key=key + ) + searches: Final = _store_searches(strict_gateway.upstream, marker) + healthy: Final = strict_gateway.gateway.request( + "POST", "/v1/rag/query", _rag_query_body(model, f"{marker} master", f"vs_{uuid.uuid4().hex}") + ) + + assert response.status_code < 500, response.text + assert searches == () + assert healthy.status_code == 200, healthy.text + + +def test_deny_by_default_searches_a_duplicated_granted_store_once(strict_gateway: StrictGateway) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 duplicate {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + body: Final[JsonObject] = { + **_rag_query_body(model, marker, store_id), + "vector_store_ids": _json_array(store_id, store_id), + } + + response: Final = strict_gateway.gateway.request("POST", "/v1/rag/query", body, key=key) + + assert response.status_code == 200, response.text + assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 + + +def test_deny_by_default_treats_null_key_grant_list_as_no_grant(strict_gateway: StrictGateway) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 null grants {uuid.uuid4().hex}" + null_grants: Final[JsonObject] = {"vector_stores": None} + key: Final = scenario.key(models=_json_array(model), object_permission=null_grants) + + response: Final = _rag_query(strict_gateway.gateway, model, marker, key, store_id=store_id) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +def _statuses_across_workers(gateway: Gateway, model: str, key: str, store_id: str) -> frozenset[int]: + return frozenset( + _rag_query(gateway, model, f"lit6035 grant change {uuid.uuid4().hex}", key, store_id=store_id).status_code + for _ in range(6) + ) + + +@pytest.mark.parametrize("scope", ("key", "team")) +def test_deny_by_default_grant_and_revoke_take_effect_on_every_worker( + strict_gateway: StrictGateway, scope: Literal["key", "team"] +) -> None: + gateway: Final = strict_gateway.gateway + with gateway.scenario() as scenario: + model: Final = _served_model(gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + key: Final = ( + scenario.key(models=_json_array(model)) + if scope == "key" + else scenario.key( + team_id=team, models=_json_array(model), object_permission=_permission_for_stores(store_id) + ) + ) + team_key: Final = scope == "team" + update_path: Final = "/team/update" if team_key else "/key/update" + identity: Final[JsonObject] = {"team_id": team} if team_key else {"key": key} + if team_key: + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) + + def statuses() -> frozenset[int]: + return _statuses_across_workers(gateway, model, key, store_id) + + before: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores(store_id)}) + granted: Final = eventually(statuses, lambda seen: seen == frozenset({200}), return_last_on_timeout=True) + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) + revoked: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) + + assert (before, granted, revoked) == (frozenset({401}), frozenset({200}), frozenset({401})) + + +def _cli_session_token(user_id: str, team_id: str) -> str: + cli_user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=[team_id], models=[]) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team_id, team_alias="cli-team") + + +@pytest.mark.parametrize("granted_by", ("team", "user")) +def test_deny_by_default_session_token_uses_only_the_resolved_team_grant( + strict_gateway: StrictGateway, granted_by: Literal["team", "user"], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt")) + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 session token {granted_by} {uuid.uuid4().hex}" + team_grants: Final = _permission_for_stores(store_id if granted_by == "team" else "vs_some_other_store") + user: Final = scenario.user( + user_role="internal_user", + object_permission=_permission_for_stores(store_id if granted_by == "user" else "vs_some_other_store"), + ) + team: Final = scenario.team( + models=_json_array(model), + object_permission=team_grants, + members_with_roles=_json_array({"role": "user", "user_id": user}), + ) + + response: Final = _rag_query( + strict_gateway.gateway, model, marker, _cli_session_token(user, team), store_id=store_id + ) + + if granted_by == "team": + assert response.status_code == 200, response.text + assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 + return + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +BurstRoute: TypeAlias = Literal["chat_retrieval_config", "search_route"] + + +def _burst_request( + gateway: Gateway, route: BurstRoute, model: str, key: str, store_id: str, marker: str +) -> httpx.Response: + if route == "search_route": + return gateway.request("POST", f"/v1/vector_stores/{store_id}/search", {"query": marker}, key=key) + return _rag_query(gateway, model, marker, key, store_id=store_id, path="/v1/chat/completions") + + +def test_deny_by_default_concurrent_burst_only_searches_granted_stores(strict_gateway: StrictGateway) -> None: + gateway: Final = strict_gateway.gateway + routes: Final[tuple[BurstRoute, ...]] = ("chat_retrieval_config", "search_route") + with gateway.scenario() as scenario: + model: Final = _served_model(gateway, scenario) + store_id: Final = CONFIG_STORE_ID + granted_key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + ungranted_key: Final = scenario.key(models=_json_array(model)) + plan: Final = tuple( + (routes[index % 2], index // 2 % 2 == 0, f"lit6035 burst {index} {uuid.uuid4().hex}") for index in range(32) + ) + upstream_observations(strict_gateway.upstream) + + with ThreadPoolExecutor(max_workers=10) as pool: + responses: Final = tuple( + pool.map( + lambda item: _burst_request( + gateway, item[0], model, granted_key if item[1] else ungranted_key, store_id, item[2] + ), + plan, + ) + ) + observed: Final = tuple( + (str(observation["path"]), str(observation["body"])) + for observation in upstream_observations(strict_gateway.upstream) + ) + + def hits(route: BurstRoute, marker: str) -> tuple[int, int]: + expected_path: Final = ( + "/chat/completions" if route == "chat_retrieval_config" else f"/vector_stores/{store_id}/search" + ) + on_route: Final = sum(path.endswith(expected_path) and marker in body for path, body in observed) + return on_route, sum(marker in body for _, body in observed) + + assert tuple(response.status_code for response in responses) == tuple( + 200 if granted else 401 for _, granted, _ in plan + ), tuple(response.text[:300] for response in responses if response.status_code not in (200, 401)) + assert tuple(hits(route, marker)[0] for route, _, marker in plan) == tuple( + 1 if granted else 0 for _, granted, _ in plan + ) + assert tuple(hits(route, marker)[1] for route, granted, marker in plan if not granted) == (0,) * 16 From 50df20c17467d0a560812d56142b1294132e2026 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 00:03:40 +0000 Subject: [PATCH 10/14] test(integration): consolidate vector_store_deny_by_default coverage to core cases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_rag_query_vector_store_allowlist.py | 373 +----------------- 1 file changed, 2 insertions(+), 371 deletions(-) diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index 3ee320994d4..d3f16e46a1e 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -1,20 +1,16 @@ from __future__ import annotations -import asyncio import json import os import time import uuid from collections.abc import Iterator, Mapping -from concurrent.futures import ThreadPoolExecutor from pathlib import Path from types import MappingProxyType from typing import Final, Literal, TypeAlias -import anthropic import httpx import jwt -import openai import pytest import yaml from cryptography.hazmat.primitives.asymmetric import rsa @@ -25,9 +21,6 @@ from integration.authorization._guardrail_opt_out import upstream_observations from pydantic import JsonValue from redis import Redis -from litellm.proxy._types import LiteLLM_UserTable -from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken - CONFIG_STORE_ID: Final = "vs_integration_config_store" PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",) @@ -73,20 +66,6 @@ def _rag_query( return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key) -def _served_model(gateway: Gateway, scenario: Scenario) -> str: - model: Final = scenario.model() - statuses: Final = eventually( - lambda: frozenset( - _rag_query(gateway, model, f"lit6035 warm {uuid.uuid4().hex}", gateway.key).status_code for _ in range(8) - ), - lambda seen: seen == frozenset({200}), - seconds=30, - return_last_on_timeout=True, - ) - assert statuses == frozenset({200}), statuses - return model - - def _searches_for_marker( gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID ) -> tuple[Mapping[str, JsonValue], ...]: @@ -195,7 +174,6 @@ def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictG environment, config=config, remove_environment=REMOVE_OPENAI_API_BASE, - workers=2, ) as gateway: yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment) @@ -205,8 +183,6 @@ StrictCase: TypeAlias = Literal[ "team_key_empty_key_grants", "team_key_empty_team_grants", "multi_store_one_ungranted", - "rag_alias_no_permission", - "chat_retrieval_config_no_permission", "jwt_user_without_grant", ] @@ -241,19 +217,6 @@ def _strict_denied_request( return strict.gateway.request( "POST", "/v1/chat/completions", body, key=partial_key ), "key_vector_store_access_denied" - if case == "rag_alias_no_permission": - alias_key: Final = scenario.key(models=models) - return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=alias_key), ( - "key_vector_store_access_denied" - ) - if case == "chat_retrieval_config_no_permission": - chat_key: Final = scenario.key(models=models) - return ( - strict.gateway.request( - "POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=chat_key - ), - "key_vector_store_access_denied", - ) user: Final = scenario.user(user_role="internal_user") return ( strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=strict.jwt(user)), @@ -268,8 +231,6 @@ def _strict_denied_request( "team_key_empty_key_grants", "team_key_empty_team_grants", "multi_store_one_ungranted", - "rag_alias_no_permission", - "chat_retrieval_config_no_permission", "jwt_user_without_grant", ), ) @@ -277,7 +238,7 @@ def test_deny_by_default_rejects_ungranted_store_before_upstream_search( strict_gateway: StrictGateway, case: StrictCase ) -> None: with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) + model: Final = scenario.model() store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}" @@ -336,7 +297,7 @@ def _strict_granted_request( ) def test_deny_by_default_searches_explicitly_granted_store(strict_gateway: StrictGateway, case: GrantedCase) -> None: with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) + model: Final = scenario.model() store_id: Final = ( CONFIG_STORE_ID if case == "standalone_key_granted_registered_store" @@ -606,333 +567,3 @@ def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: Stric strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()}) assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401 - - -def _store_searches(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: - return tuple( - observation - for observation in upstream_observations(gateway) - if str(observation["path"]).startswith("/vector_stores/") and marker in str(observation["body"]) - ) - - -def _marker_observations(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: - return tuple(observation for observation in upstream_observations(gateway) if marker in str(observation["body"])) - - -def _messages_body(model: str, marker: str, store_id: str) -> JsonObject: - body: Final[JsonObject] = { - "model": model, - "max_tokens": 16, - "messages": _json_array({"role": "user", "content": marker}), - "vector_store_ids": _json_array(store_id), - } - return body - - -def _file_search_tools(store_id: JsonValue) -> JsonValue: - return _json_array({"type": "file_search", "vector_store_ids": store_id}) - - -RequestShape: TypeAlias = Literal["messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids"] - - -def _shape_request(shape: RequestShape, model: str, marker: str, store_id: str) -> tuple[str, JsonObject]: - messages: Final = _json_array({"role": "user", "content": marker}) - if shape == "messages_vector_store_ids": - return "/v1/messages", _messages_body(model, marker, store_id) - if shape == "chat_stream_file_search": - stream_body: Final[JsonObject] = { - "model": model, - "messages": messages, - "stream": True, - "tools": _file_search_tools(_json_array(store_id)), - } - return "/v1/chat/completions", stream_body - chat_body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": _json_array(store_id)} - return "/v1/chat/completions", chat_body - - -@pytest.mark.parametrize("shape", ("messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids")) -def test_deny_by_default_covers_messages_streaming_and_top_level_store_ids( - strict_gateway: StrictGateway, shape: RequestShape -) -> None: - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - marker: Final = f"lit6035 shape {shape} {uuid.uuid4().hex}" - key: Final = scenario.key(models=_json_array(model)) - path, body = _shape_request(shape, model, marker, store_id) - - response: Final = strict_gateway.gateway.request("POST", path, body, key=key) - - assert response.status_code == 401, response.text - assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text - assert _marker_observations(strict_gateway.upstream, marker) == () - - -SdkClient: TypeAlias = Literal["openai_sync", "openai_async", "anthropic_sync"] - - -def _sdk_denial(gateway: Gateway, client: SdkClient, key: str, model: str, marker: str, store_id: str) -> int: - base_url: Final = str(gateway.client.base_url) - if client == "anthropic_sync": - with pytest.raises(anthropic.AuthenticationError) as anthropic_denied: - anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create( - model=model, - max_tokens=16, - messages=[{"role": "user", "content": marker}], - extra_body={"vector_store_ids": _json_array(store_id)}, - ) - return anthropic_denied.value.status_code - tools: Final = [{"type": "file_search", "vector_store_ids": [store_id]}] - if client == "openai_sync": - with pytest.raises(openai.AuthenticationError) as sync_denied: - openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0).responses.create( - model=model, input=marker, tools=tools - ) - return sync_denied.value.status_code - - async def create() -> None: - async with openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as async_client: - await async_client.responses.create(model=model, input=marker, tools=tools) - - with pytest.raises(openai.AuthenticationError) as async_denied: - asyncio.run(create()) - return async_denied.value.status_code - - -@pytest.mark.parametrize("client", ("openai_sync", "openai_async", "anthropic_sync")) -def test_deny_by_default_rejects_sdk_clients_without_grant(strict_gateway: StrictGateway, client: SdkClient) -> None: - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - marker: Final = f"lit6035 sdk {client} {uuid.uuid4().hex}" - key: Final = scenario.key(models=_json_array(model)) - - status: Final = _sdk_denial(strict_gateway.gateway, client, key, model, marker, store_id) - - assert status == 401 - assert _marker_observations(strict_gateway.upstream, marker) == () - - -MalformedShape: TypeAlias = Literal[ - "ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id" -] - - -def _malformed_body(shape: MalformedShape, model: str, marker: str) -> JsonObject: - messages: Final = _json_array({"role": "user", "content": marker}) - if shape == "tool_ids_string": - tool_ids_body: Final[JsonObject] = { - "model": model, - "messages": messages, - "tools": _file_search_tools(f"vs_{uuid.uuid4().hex}"), - } - return tool_ids_body - if shape == "tool_is_string": - string_tool_body: Final[JsonObject] = { - "model": model, - "messages": messages, - "tools": _json_array("file_search"), - } - return string_tool_body - ids: Final[Mapping[MalformedShape, JsonValue]] = MappingProxyType( - { - "ids_string": f"vs_{uuid.uuid4().hex}", - "ids_int": _json_array(123), - "ids_empty_string": _json_array(""), - "oversized_id": _json_array("vs_" + "x" * 5000), - } - ) - body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": ids[shape]} - return body - - -@pytest.mark.parametrize( - "shape", ("ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id") -) -def test_deny_by_default_malformed_store_ids_never_search_a_store( - strict_gateway: StrictGateway, shape: MalformedShape -) -> None: - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - marker: Final = f"lit6035 malformed {shape} {uuid.uuid4().hex}" - key: Final = scenario.key(models=_json_array(model)) - - response: Final = strict_gateway.gateway.request( - "POST", "/v1/chat/completions", _malformed_body(shape, model, marker), key=key - ) - searches: Final = _store_searches(strict_gateway.upstream, marker) - healthy: Final = strict_gateway.gateway.request( - "POST", "/v1/rag/query", _rag_query_body(model, f"{marker} master", f"vs_{uuid.uuid4().hex}") - ) - - assert response.status_code < 500, response.text - assert searches == () - assert healthy.status_code == 200, healthy.text - - -def test_deny_by_default_searches_a_duplicated_granted_store_once(strict_gateway: StrictGateway) -> None: - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - marker: Final = f"lit6035 duplicate {uuid.uuid4().hex}" - key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) - body: Final[JsonObject] = { - **_rag_query_body(model, marker, store_id), - "vector_store_ids": _json_array(store_id, store_id), - } - - response: Final = strict_gateway.gateway.request("POST", "/v1/rag/query", body, key=key) - - assert response.status_code == 200, response.text - assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 - - -def test_deny_by_default_treats_null_key_grant_list_as_no_grant(strict_gateway: StrictGateway) -> None: - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - marker: Final = f"lit6035 null grants {uuid.uuid4().hex}" - null_grants: Final[JsonObject] = {"vector_stores": None} - key: Final = scenario.key(models=_json_array(model), object_permission=null_grants) - - response: Final = _rag_query(strict_gateway.gateway, model, marker, key, store_id=store_id) - - assert response.status_code == 401, response.text - assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text - assert _marker_observations(strict_gateway.upstream, marker) == () - - -def _statuses_across_workers(gateway: Gateway, model: str, key: str, store_id: str) -> frozenset[int]: - return frozenset( - _rag_query(gateway, model, f"lit6035 grant change {uuid.uuid4().hex}", key, store_id=store_id).status_code - for _ in range(6) - ) - - -@pytest.mark.parametrize("scope", ("key", "team")) -def test_deny_by_default_grant_and_revoke_take_effect_on_every_worker( - strict_gateway: StrictGateway, scope: Literal["key", "team"] -) -> None: - gateway: Final = strict_gateway.gateway - with gateway.scenario() as scenario: - model: Final = _served_model(gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id)) - key: Final = ( - scenario.key(models=_json_array(model)) - if scope == "key" - else scenario.key( - team_id=team, models=_json_array(model), object_permission=_permission_for_stores(store_id) - ) - ) - team_key: Final = scope == "team" - update_path: Final = "/team/update" if team_key else "/key/update" - identity: Final[JsonObject] = {"team_id": team} if team_key else {"key": key} - if team_key: - gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) - - def statuses() -> frozenset[int]: - return _statuses_across_workers(gateway, model, key, store_id) - - before: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) - gateway.post(update_path, {**identity, "object_permission": _permission_for_stores(store_id)}) - granted: Final = eventually(statuses, lambda seen: seen == frozenset({200}), return_last_on_timeout=True) - gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) - revoked: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) - - assert (before, granted, revoked) == (frozenset({401}), frozenset({200}), frozenset({401})) - - -def _cli_session_token(user_id: str, team_id: str) -> str: - cli_user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=[team_id], models=[]) - return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team_id, team_alias="cli-team") - - -@pytest.mark.parametrize("granted_by", ("team", "user")) -def test_deny_by_default_session_token_uses_only_the_resolved_team_grant( - strict_gateway: StrictGateway, granted_by: Literal["team", "user"], monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt")) - with strict_gateway.gateway.scenario() as scenario: - model: Final = _served_model(strict_gateway.gateway, scenario) - store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" - marker: Final = f"lit6035 session token {granted_by} {uuid.uuid4().hex}" - team_grants: Final = _permission_for_stores(store_id if granted_by == "team" else "vs_some_other_store") - user: Final = scenario.user( - user_role="internal_user", - object_permission=_permission_for_stores(store_id if granted_by == "user" else "vs_some_other_store"), - ) - team: Final = scenario.team( - models=_json_array(model), - object_permission=team_grants, - members_with_roles=_json_array({"role": "user", "user_id": user}), - ) - - response: Final = _rag_query( - strict_gateway.gateway, model, marker, _cli_session_token(user, team), store_id=store_id - ) - - if granted_by == "team": - assert response.status_code == 200, response.text - assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 - return - assert response.status_code == 401, response.text - assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text - assert _marker_observations(strict_gateway.upstream, marker) == () - - -BurstRoute: TypeAlias = Literal["chat_retrieval_config", "search_route"] - - -def _burst_request( - gateway: Gateway, route: BurstRoute, model: str, key: str, store_id: str, marker: str -) -> httpx.Response: - if route == "search_route": - return gateway.request("POST", f"/v1/vector_stores/{store_id}/search", {"query": marker}, key=key) - return _rag_query(gateway, model, marker, key, store_id=store_id, path="/v1/chat/completions") - - -def test_deny_by_default_concurrent_burst_only_searches_granted_stores(strict_gateway: StrictGateway) -> None: - gateway: Final = strict_gateway.gateway - routes: Final[tuple[BurstRoute, ...]] = ("chat_retrieval_config", "search_route") - with gateway.scenario() as scenario: - model: Final = _served_model(gateway, scenario) - store_id: Final = CONFIG_STORE_ID - granted_key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) - ungranted_key: Final = scenario.key(models=_json_array(model)) - plan: Final = tuple( - (routes[index % 2], index // 2 % 2 == 0, f"lit6035 burst {index} {uuid.uuid4().hex}") for index in range(32) - ) - upstream_observations(strict_gateway.upstream) - - with ThreadPoolExecutor(max_workers=10) as pool: - responses: Final = tuple( - pool.map( - lambda item: _burst_request( - gateway, item[0], model, granted_key if item[1] else ungranted_key, store_id, item[2] - ), - plan, - ) - ) - observed: Final = tuple( - (str(observation["path"]), str(observation["body"])) - for observation in upstream_observations(strict_gateway.upstream) - ) - - def hits(route: BurstRoute, marker: str) -> tuple[int, int]: - expected_path: Final = ( - "/chat/completions" if route == "chat_retrieval_config" else f"/vector_stores/{store_id}/search" - ) - on_route: Final = sum(path.endswith(expected_path) and marker in body for path, body in observed) - return on_route, sum(marker in body for _, body in observed) - - assert tuple(response.status_code for response in responses) == tuple( - 200 if granted else 401 for _, granted, _ in plan - ), tuple(response.text[:300] for response in responses if response.status_code not in (200, 401)) - assert tuple(hits(route, marker)[0] for route, _, marker in plan) == tuple( - 1 if granted else 0 for _, granted, _ in plan - ) - assert tuple(hits(route, marker)[1] for route, granted, marker in plan if not granted) == (0,) * 16 From 3e5fe50faa7d2674acf40d1a20a1ee7b8f2b148e Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 01:21:44 +0000 Subject: [PATCH 11/14] fix(proxy): reject invalid vector_store_deny_by_default at config load and return 400 for malformed vector_store_ids Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 102 ++++++++++++++---- litellm/proxy/proxy_server.py | 8 ++ ...st_auth_checks_object_access_and_lookup.py | 72 +++++++++++++ .../proxy/proxy_server/test_proxy_config.py | 17 +++ 4 files changed, 178 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 30f8349b79d..8b42b243cea 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -15,6 +15,7 @@ import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from functools import partial +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias @@ -146,7 +147,6 @@ from litellm.router import Router from litellm.types.proxy.auth.auth_checks import UserNotFoundError from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget from litellm.utils import get_utc_datetime -from litellm.vector_stores.vector_store_registry import VectorStoreRegistry from .auth_checks_organization import ( add_team_org_context_to_request_body, @@ -1286,15 +1286,7 @@ async def common_checks( team_object=team_object, valid_token=valid_token, user_object=user_object, - deny_by_default=ConfigGeneralSettings.model_validate( - MappingProxyType( - { - "vector_store_deny_by_default": _typed_request_body(general_settings).get( - "vector_store_deny_by_default", False - ) - } - ) - ).vector_store_deny_by_default, + deny_by_default=_vector_store_deny_by_default(_typed_request_body(general_settings)), ) # 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path) @@ -6644,6 +6636,73 @@ def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool ) +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 + general_settings another way enables the policy, so only vector store requests are denied. + """ + try: + return ConfigGeneralSettings.model_validate( + MappingProxyType( + {"vector_store_deny_by_default": general_settings.get("vector_store_deny_by_default", False)} + ) + ).vector_store_deny_by_default + except ValidationError: + return True + + +_VECTOR_STORE_IDS_ADAPTER: Final[TypeAdapter[list[str]]] = TypeAdapter(list[str]) +_TOOLS_ADAPTER: Final[TypeAdapter[list[object]]] = TypeAdapter(list[object]) +_TOOL_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) + + +def _validated_vector_store_ids(value: object) -> tuple[str, ...]: + if value is None: + return () + try: + return tuple(_VECTOR_STORE_IDS_ADAPTER.validate_python(value, strict=True)) + except ValidationError: + raise _malformed_vector_store_ids() from None + + +def _malformed_vector_store_ids() -> ProxyException: + return ProxyException( + message="vector_store_ids must be a list of strings", + type="invalid_request_error", + param="vector_store_ids", + code=status.HTTP_400_BAD_REQUEST, + ) + + +def _tools(tools: object) -> tuple[object, ...]: + try: + return tuple(_TOOLS_ADAPTER.validate_python(tools, strict=True)) + except ValidationError: + return () + + +def _tool_vector_store_ids(tool: object) -> tuple[str, ...]: + try: + tool_fields: Final = _TOOL_ADAPTER.validate_python(tool, strict=True) + except ValidationError: + return () + return _validated_vector_store_ids(tool_fields.get("vector_store_ids")) + + +def _strict_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]: + """ + Same fields VectorStoreRegistry.get_vector_store_ids_to_run reads, but a vector_store_ids that is + not a list of strings is a 400 instead of being skipped or iterated, and tools that are not + objects name no store. + """ + return tuple( + chain( + _validated_vector_store_ids(request_body.get("vector_store_ids")), + chain.from_iterable(_tool_vector_store_ids(tool) for tool in _tools(request_body.get("tools"))), + ) + ) + + def _require_vector_store_grant( object_type: Literal["key", "team", "user"], vector_store_ids_to_run: Sequence[str], @@ -6694,7 +6753,8 @@ async def vector_store_access_check( rescue nor restrict a key or team request. Raises ProxyException (401, `{key,team,user}_vector_store_access_denied`) on the first identity - that does not grant a requested store. + that does not grant a requested store, and with the flag on, ProxyException (400, + `invalid_request_error`) when a `vector_store_ids` field is not a list of strings. """ from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -6705,18 +6765,18 @@ async def vector_store_access_check( verbose_proxy_logger.debug("Prisma client not found, skipping vector store access check") return True - vector_store_registry: Final = ( - VectorStoreRegistry() - if litellm.vector_store_registry is None and deny_by_default - else litellm.vector_store_registry - ) registry_ids: Final = ( - vector_store_registry.get_vector_store_ids_to_run( - non_default_params=request_body, tools=request_body.get("tools", None) + _strict_requested_vector_store_ids(_typed_request_body(request_body)) + if deny_by_default + else ( + litellm.vector_store_registry.get_vector_store_ids_to_run( + non_default_params=request_body, tools=request_body.get("tools", None) + ) + if litellm.vector_store_registry is not None + else None ) - if vector_store_registry is not None - else None - ) or () + or () + ) rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body)) rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else () vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids))) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e406e6faec4..f308e6b9d02 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6657,6 +6657,14 @@ class ProxyConfig: if general_settings is None: general_settings = {} + typed_general_settings: Final = _GENERAL_SETTINGS_VIEW.validate_python(general_settings) + if "vector_store_deny_by_default" in typed_general_settings: + ConfigGeneralSettings.model_validate( + MappingProxyType( + {"vector_store_deny_by_default": typed_general_settings["vector_store_deny_by_default"]} + ) + ) + if general_settings.get("mcp_advertised_versions") is not None: from litellm.types.mcp import MCPAdvertisedVersions 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 03164f0e4d7..ceb9111a68f 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 @@ -2275,6 +2275,78 @@ async def test_deny_by_default_reads_requested_vector_stores_without_a_registry( ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("request_body", "key_vector_stores", "denied_by"), + [ + ({"tools": [1]}, None, None), + ({"tools": 1}, None, None), + ({"tools": [{"type": "file_search", "vector_store_ids": None}]}, None, None), + ({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREA"], None), + ({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREB"], _KEY_DENIED), + ], + ids=["int-tool", "non-list-tools", "null-tool-ids", "int-tool-beside-granted", "int-tool-beside-ungranted"], +) +async def test_deny_by_default_ignores_tools_that_name_no_vector_store( + request_body: Mapping[str, object], key_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None +): + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + _virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission"), + None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores), + request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body}, + ), + denied_by, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "request_body", + [ + {"tools": [{"type": "file_search", "vector_store_ids": "KBSTOREA"}]}, + {"tools": [{"type": "file_search", "vector_store_ids": [1]}]}, + {"vector_store_ids": "KBSTOREA"}, + ], + ids=["string-tool-ids", "int-tool-id", "string-top-level-ids"], +) +async def test_deny_by_default_rejects_malformed_vector_store_ids_as_bad_request(request_body: Mapping[str, object]): + with pytest.raises(ProxyException) as exc_info: + await _common_checks_for_rag_query( + {"vector_store_deny_by_default": True}, + _virtual_key(object_permission_id="key-permission"), + SimpleNamespace(vector_stores=["KBSTOREA"]), + request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body}, + ) + assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == ( + "invalid_request_error", + "vector_store_ids", + "400", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("request_body", "denied_by"), + [ + ({"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]}, None), + (None, _KEY_DENIED), + ], + ids=["no-vector-store", "ungranted-vector-store"], +) +@pytest.mark.parametrize("flag_value", [None, "enabled"], ids=["null", "string"]) +async def test_invalid_deny_by_default_value_only_denies_vector_store_requests( + flag_value: object, request_body: Mapping[str, object] | None, denied_by: ProxyErrorTypes | None +): + await _assert_rag_query_outcome( + _common_checks_for_rag_query( + {"vector_store_deny_by_default": flag_value}, _virtual_key(), None, request_body=request_body + ), + denied_by, + ) + + @pytest.mark.asyncio async def test_keyless_user_grant_is_read_through_the_object_permission_cache(): cache: Final = UserApiKeyCache(default_in_memory_ttl=60) diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 804a372a1c9..0bd4986a1a6 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -2293,6 +2293,23 @@ async def test_load_config_yaml_vector_store_deny_by_default_is_boolean( assert ConfigGeneralSettings.model_validate(dict(general_settings)).vector_store_deny_by_default is expected +@pytest.mark.asyncio +@pytest.mark.parametrize("yaml_value", ["", "enabled"], ids=["null", "string"]) +async def test_load_config_rejects_non_boolean_vector_store_deny_by_default( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str +): + config_file: Final = tmp_path / "vector_store.yaml" + config_file.write_text( + f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n" + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + with pytest.raises(ValidationError, match="vector_store_deny_by_default"): + await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + + @pytest.mark.asyncio async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch): """Regression: router_settings.plugins dotted-path strings must be resolved to From 64158253a618c7a9e8fe5c96d46878e46cf9f60b Mon Sep 17 00:00:00 2001 From: mrinal Date: Sun, 4 Oct 2026 08:45:53 +0000 Subject: [PATCH 12/14] 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"] From bdc00ac85498554ac70b6507d09eeaf1309224a8 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sun, 4 Oct 2026 09:29:51 +0000 Subject: [PATCH 13/14] test(auth): return typed permission rows in request flow vector store tests Co-authored-by: mrinal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/auth/test_user_api_key_auth_request_flow.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 76f7fb2aae6..df31595f7ec 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -28,6 +28,7 @@ from litellm.proxy._types import ( LiteLLM_JWTAuth, LiteLLM_BudgetTable, LiteLLM_EndUserTable, + LiteLLM_ObjectPermissionTable, LiteLLM_OrganizationTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, @@ -5328,13 +5329,16 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( request._url = URL(url="/v1/rag/query") database: Final = MagicMock() database.db.litellm_objectpermissiontable.find_unique = AsyncMock( - side_effect=lambda where: SimpleNamespace(vector_stores=["KBSTOREA"]) + side_effect=lambda where: LiteLLM_ObjectPermissionTable( + object_permission_id=where["object_permission_id"], vector_stores=["KBSTOREA"] + ) if where["object_permission_id"] == "key-permission" else None ) attrs: Final = { **_proxy_attrs_for_centralized_checks(), "prisma_client": database, + "proxy_logging_obj": MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock())), "general_settings": {"vector_store_deny_by_default": deny_by_default}, } for name, value in attrs.items(): @@ -5398,11 +5402,14 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te } database: Final = MagicMock() database.db.litellm_objectpermissiontable.find_unique = AsyncMock( - side_effect=lambda where: SimpleNamespace(vector_stores=grants[where["object_permission_id"]]) + side_effect=lambda where: LiteLLM_ObjectPermissionTable( + object_permission_id=where["object_permission_id"], vector_stores=grants[where["object_permission_id"]] + ) ) attrs: Final = { **_proxy_attrs_for_centralized_checks(), "prisma_client": database, + "proxy_logging_obj": MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock())), "general_settings": {"vector_store_deny_by_default": True}, } for name, value in attrs.items(): From ee9cb4ea985c38087187c14610ab83adba494252 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sun, 4 Oct 2026 10:02:02 +0000 Subject: [PATCH 14/14] refactor(auth): validate vector store ids and tools as immutable sequences Co-authored-by: mrinal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 90b4e6ca2a3..d6b92ec18ac 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -6736,8 +6736,8 @@ def _vector_store_deny_by_default(general_settings: Mapping[str, object]) -> boo return True -_VECTOR_STORE_IDS_ADAPTER: Final[TypeAdapter[list[str]]] = TypeAdapter(list[str]) -_TOOLS_ADAPTER: Final[TypeAdapter[list[object]]] = TypeAdapter(list[object]) +_VECTOR_STORE_IDS_ADAPTER: Final[TypeAdapter[Sequence[str]]] = TypeAdapter(Sequence[str]) +_TOOLS_ADAPTER: Final[TypeAdapter[Sequence[object]]] = TypeAdapter(Sequence[object]) _TOOL_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])