diff --git a/litellm/litellm_core_utils/error_normalization.py b/litellm/litellm_core_utils/error_normalization.py index 726fdcd1540..4426d42f767 100644 --- a/litellm/litellm_core_utils/error_normalization.py +++ b/litellm/litellm_core_utils/error_normalization.py @@ -107,6 +107,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 c1a97aa73c4..23cd64b8a49 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3022,6 +3022,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 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, 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.", @@ -4512,6 +4516,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 +4555,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 +4566,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 460b25ddac7..3ebd1a2cfe0 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 @@ -34,6 +35,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,7 +46,9 @@ 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, LiteLLM_BudgetTable, LiteLLM_EndUserTable, @@ -313,12 +317,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 @@ -1353,6 +1351,8 @@ async def common_checks( request_body=request_body, team_object=team_object, valid_token=valid_token, + user_object=user_object, + 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) @@ -6792,17 +6792,210 @@ 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 +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 + and valid_token.team_id != UI_TEAM_ID + ) + + +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 + 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[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]) + + +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: GrantLayer, + vector_store_ids_to_run: Sequence[str], + object_permission: _VectorStorePermissionsRow | None, +) -> None: + if object_permission is None or not object_permission.vector_stores: + raise ProxyException( + 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=object_type, + vector_store_ids_to_run=vector_store_ids_to_run, + object_permissions=object_permission, + ) + + +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, valid_token: UserAPIKeyAuth | None, + *, + user_object: LiteLLM_UserTable | None = None, + 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, 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 + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache ######################################################### # Get the vector store the user is trying to access @@ -6812,12 +7005,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) + _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 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))) @@ -6828,38 +7026,28 @@ 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( - where={"object_permission_id": valid_token.object_permission_id}, - ) - if key_object_permission is not 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, + ) + 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="key", + object_type=layer, 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: - 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, + object_permissions=grant, ) 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/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index adb6b86321d..65f9b09a4d3 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3063,6 +3063,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/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 69efb9c209f..98df4ae46b1 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1458,11 +1458,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/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 818c0590707..aeb317c8905 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -165,6 +165,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, @@ -2546,6 +2547,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/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d1147c48558..6d522759eb7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6713,6 +6713,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/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/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index 896c88c68bb..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,5 +1,8 @@ from __future__ import annotations +import json +import os +import time import uuid from collections.abc import Iterator, Mapping from pathlib import Path @@ -7,16 +10,22 @@ from types import MappingProxyType from typing import Final, Literal, TypeAlias import httpx +import jwt import pytest import yaml -from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value +from cryptography.hazmat.primitives.asymmetric import rsa +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] @@ -99,6 +108,211 @@ 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, + 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] = { + "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") + config: Final = _strict_config(directory) + environment: Final = MappingProxyType({**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url}) + with owned_proxy( + upstream_gateway, + directory, + environment, + config=config, + remove_environment=REMOVE_OPENAI_API_BASE, + ) as gateway: + yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment) + + +StrictCase: TypeAlias = Literal[ + "standalone_key_no_permission", + "team_key_empty_key_grants", + "team_key_empty_team_grants", + "multi_store_one_ungranted", + "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": + 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: 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=team_key + ), (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] = { + **_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=partial_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", + "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"): + 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": + 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)) + 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 +434,136 @@ 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) == () + + +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 + + +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/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 9ce23c442f6..e855d8cd346 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, @@ -90,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, ) @@ -1646,8 +1649,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() @@ -1688,12 +1692,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() @@ -1755,12 +1759,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 ( @@ -1786,6 +1790,592 @@ 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 + + +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: Final = 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 + + +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, + key_permission: SimpleNamespace | None, + team_object: LiteLLM_TeamTable | None = None, + 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, + user_api_key_cache: UserApiKeyCache | None = None, + vector_store_registry: VectorStoreRegistry | None = None, +) -> bool: + permissions: Final = { + "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( + side_effect=lambda where: permissions.get(where["object_permission_id"]) + ) + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + 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", 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, + team_object=team_object, + user_object=user_object, + 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], 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) == (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", "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], denied_by: ProxyErrorTypes | None +): + 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", "denied_by"), + [ + (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, 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=vector_stores), + ), + denied_by, + ) + + +@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: 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 + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("general_settings", "vector_stores", "denied_by"), + [ + ({}, None, None), + ({"vector_store_deny_by_default": False}, None, None), + ({"vector_store_deny_by_default": False}, ["KBSTOREB"], _KEY_DENIED), + ], + ids=["flag-omitted-no-record", "flag-false-no-record", "flag-false-excludes"], +) +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: Final = _virtual_key( + object_permission_id=None if vector_stores is None else "key-permission", user_role=LitellmUserRoles.PROXY_ADMIN + ) + 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 + ) + + +@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( + _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: 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 + ) + + +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: Final = 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, + ) + + +@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, + ) + + +@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 +@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) + 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, + ) + + +@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 @@ -10346,7 +10936,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) @@ -10498,9 +11090,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") ) @@ -10523,7 +11118,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/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index d4d11527994..5de1d9300fe 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,11 +22,13 @@ 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, LiteLLM_BudgetTable, LiteLLM_EndUserTable, + LiteLLM_ObjectPermissionTable, LiteLLM_OrganizationTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, @@ -5304,6 +5307,222 @@ 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: 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: 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: 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(): + 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: 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.team_vector_store_access_denied, + "vector_store", + "401", + ) + + +@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: Final = UserAPIKeyAuth( + user_id="user-1", + team_id="team-1", + team_models=["gpt-4o-mini"], + user_role=LitellmUserRoles.INTERNAL_USER, + ) + request: Final = Request(scope={"type": "http"}) + request._url = URL(url="/v1/rag/query") + grants: Final = { + "team-1-grants-a": ["KBSTOREA"], + "team-1-grants-none": [], + "team-2-grants-a": ["KBSTOREA"], + "user-grants-a": ["KBSTOREA"], + } + database: Final = MagicMock() + database.db.litellm_objectpermissiontable.find_unique = AsyncMock( + 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(): + 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: 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.team_vector_store_access_denied, + "vector_store", + "401", + ) + + +@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/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 7150ba892df..912622da4bd 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -4066,6 +4066,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. diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 6f6dbe1e954..6c5447fce64 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -8417,6 +8417,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 @@ -8570,6 +8571,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"] diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 2f5eb81bec8..b37576dab80 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, @@ -2329,6 +2329,42 @@ 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: 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) + + _, _, 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 +@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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c5b8c0e4d14..c632ad01203 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -30083,6 +30083,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: {