From 3e5fe50faa7d2674acf40d1a20a1ee7b8f2b148e Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 01:21:44 +0000 Subject: [PATCH] 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