From a38fff65601ce67b8ee2c3a38dc14eacdfa646e0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 13:48:41 -0700 Subject: [PATCH] fix(proxy): enforce key/team vector_stores allowlist on /v1/rag/query (#43953) * add test case for /rag/query and stronger auth check * style(proxy): ruff format auth_checks.py Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): build rag query vector store ids immutably and test the no-registry path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): integration coverage for /v1/rag/query vector store allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): audit cells for /v1/rag/query vector store allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): consolidate vector store allowlist audit coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): type RAG vector store request body Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Mrinal Chanshetty Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 32 ++- .../test_rag_query_vector_store_allowlist.py | 222 ++++++++++++++++++ ...st_auth_checks_object_access_and_lookup.py | 60 +++++ 3 files changed, 306 insertions(+), 8 deletions(-) create mode 100644 tests/integration/authorization/test_rag_query_vector_store_allowlist.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 4b7b8290a30..dbd6f28a183 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -6609,6 +6609,19 @@ def _is_wildcard_pattern(allowed_model_pattern: str) -> bool: return "*" in allowed_model_pattern +def _get_rag_query_vector_store_id(request_body: Mapping[str, object]) -> str | None: + """ + /v1/rag/query carries its vector store in retrieval_config.vector_store_id, + not in vector_store_ids or tools[].vector_store_ids. + """ + retrieval_config: Final = request_body.get("retrieval_config") + if not isinstance(retrieval_config, dict): + return None + + vector_store_id: Final = retrieval_config.get("vector_store_id") + return vector_store_id if isinstance(vector_store_id, str) and vector_store_id else None + + async def vector_store_access_check( request_body: dict, team_object: LiteLLM_TeamTable | None, @@ -6628,13 +6641,16 @@ async def vector_store_access_check( verbose_proxy_logger.debug("Prisma client not found, skipping vector store access check") return True - if litellm.vector_store_registry is None: - verbose_proxy_logger.debug("Vector store registry not found, skipping vector store access check") - return True - - vector_store_ids_to_run: Final = litellm.vector_store_registry.get_vector_store_ids_to_run( - non_default_params=request_body, tools=request_body.get("tools", None) - ) + registry_ids: Final = ( + 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 () + 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))) if not vector_store_ids_to_run: verbose_proxy_logger.debug("Vector store to run not found, skipping vector store access check") return True @@ -6674,7 +6690,7 @@ async def vector_store_access_check( def _can_object_call_vector_stores( object_type: Literal["key", "team", "org"], - vector_store_ids_to_run: list[str], + vector_store_ids_to_run: Sequence[str], object_permissions: _VectorStorePermissionsRow | 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 new file mode 100644 index 00000000000..896c88c68bb --- /dev/null +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +import uuid +from collections.abc import Iterator, Mapping +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value +from integration._support.process import owned_proxy +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",) +JsonObject: TypeAlias = dict[str, JsonValue] + + +def _json_array(*values: JsonValue) -> JsonValue: + return [*values] # mutable-ok: request payloads and YAML sequences require list values + + +def _permission_for_stores(*store_ids: str) -> JsonObject: + permission: Final[JsonObject] = {"vector_stores": _json_array(*store_ids)} + return permission + + +def _key_for_scope(scenario: Scenario, model: str, scope: Literal["key", "team"], store_id: str) -> str: + if scope == "key": + return scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + return scenario.key(team_id=team, models=_json_array(model)) + + +def _rag_query_body(model: str, marker: str, store_id: str) -> JsonObject: + body: Final[JsonObject] = { + "model": model, + "messages": _json_array({"role": "user", "content": marker}), + "retrieval_config": {"vector_store_id": store_id, "custom_llm_provider": "openai", "top_k": 1}, + } + return body + + +def _rag_query( + gateway: Gateway, + model: str, + marker: str, + key: str, + *, + store_id: str = CONFIG_STORE_ID, + path: str = "/v1/rag/query", +) -> httpx.Response: + return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key) + + +def _searches_for_marker( + gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID +) -> tuple[Mapping[str, JsonValue], ...]: + search_path: Final = f"/vector_stores/{store_id}/search" + return tuple( + observation + for observation in upstream_observations(gateway) + if observation["path"] == search_path and marker in str(observation["body"]) + ) + + +def _no_registry_config(directory: Path) -> Path: + config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text())) + config_without_registry: Final[Mapping[str, JsonValue]] = MappingProxyType( + {name: value for name, value in config.items() if name != "vector_store_registry"} + ) + yaml_config: Final[JsonObject] = {**config_without_registry, "model_list": _json_array()} + path: Final = directory / "proxy_no_vector_store_registry.yaml" + path.write_text(yaml.safe_dump(yaml_config)) + return path + + +def _openai_environment(gateway: Gateway) -> Mapping[str, str]: + return MappingProxyType({"OPENAI_BASE_URL": gateway.upstream_url, "OPENAI_API_KEY": "synthetic-openai-key"}) + + +@pytest.fixture(scope="module") +def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Gateway]]: + with gateway_from_environment() as upstream_gateway: + directory: Final = tmp_path_factory.mktemp("rag_query_no_registry") + config: Final = _no_registry_config(directory) + with owned_proxy( + upstream_gateway, + directory, + _openai_environment(upstream_gateway), + config=config, + remove_environment=REMOVE_OPENAI_API_BASE, + workers=2, + ) as no_registry_gateway: + yield no_registry_gateway, upstream_gateway + + +@pytest.mark.parametrize( + ("scope", "error_type"), + (("key", "key_vector_store_access_denied"), ("team", "team_vector_store_access_denied")), +) +def test_rag_query_is_denied_when_key_or_team_allowlist_excludes_store( + gateway: Gateway, scope: Literal["key", "team"], error_type: str +) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, scope, "vs_some_other_store") + marker: Final = f"lit5610 rag query denied {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == error_type, response.text + assert _searches_for_marker(gateway, marker) == () + + +@pytest.mark.parametrize("scope", ("key", "team")) +def test_rag_query_searches_configured_store_when_allowlist_includes_it( + gateway: Gateway, scope: Literal["key", "team"] +) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, scope, CONFIG_STORE_ID) + marker: Final = f"lit5610 rag query allowed {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(gateway, marker) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +def test_rag_query_without_key_object_permission_can_search_store(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=_json_array(model)) + marker: Final = f"lit5610 rag query no permission {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(gateway, marker) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +@pytest.mark.parametrize("scope", ("team", "key")) +def test_no_registry_rag_query_denies_unregistered_store_when_allowlist_excludes( + no_registry_gateways: tuple[Gateway, Gateway], scope: Literal["team", "key"] +) -> None: + no_registry_gateway, upstream_gateway = no_registry_gateways + with no_registry_gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + key: Final = _key_for_scope(scenario, model, scope, "vs_some_other_store") + marker: Final = f"lit5610 no registry denied {scope} {uuid.uuid4().hex}" + error_type: Final = "team_vector_store_access_denied" if scope == "team" else "key_vector_store_access_denied" + + response: Final = _rag_query(no_registry_gateway, model, marker, key, store_id=store_id) + searches: Final = _searches_for_marker(upstream_gateway, marker, store_id) + + assert response.status_code == 401, f"{response.text}; scripted_upstream_searches={searches!r}" + assert response.json()["error"]["type"] == error_type, response.text + assert searches == () + + +def test_no_registry_rag_query_allows_team_allowlisted_unregistered_store( + no_registry_gateways: tuple[Gateway, Gateway], +) -> None: + no_registry_gateway, upstream_gateway = no_registry_gateways + with no_registry_gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + key: Final = _key_for_scope(scenario, model, "team", store_id) + marker: Final = f"lit5610 no registry allowed {uuid.uuid4().hex}" + + response: Final = _rag_query(no_registry_gateway, model, marker, key, store_id=store_id) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(upstream_gateway, marker, store_id) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +def test_chat_completions_top_level_retrieval_config_uses_team_allowlist(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, "team", "vs_some_other_store") + marker: Final = f"lit5610 chat top-level retrieval config denied {uuid.uuid4().hex}" + body: Final[JsonObject] = { + "model": model, + "messages": _json_array({"role": "user", "content": marker}), + "retrieval_config": { + "vector_store_id": CONFIG_STORE_ID, + "custom_llm_provider": "openai", + "top_k": 1, + }, + } + + response: Final = gateway.request("POST", "/v1/chat/completions", body, key=key) + observations: Final = upstream_observations(gateway) + + assert response.status_code == 401, f"{response.text}; scripted_upstream_observations={observations!r}" + assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text + + +def test_rag_query_alias_denies_store_when_team_allowlist_excludes(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, "team", "vs_some_other_store") + marker: Final = f"lit5610 rag query alias denied {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key, path="/rag/query") + + 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) == () 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 353249dddf0..6c8b6571991 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 @@ -94,6 +94,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( tag_registry_cache_key, ) from litellm.utils import get_utc_datetime +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry def _rendered_log_message(call): @@ -1753,6 +1754,65 @@ async def test_vector_store_access_check_with_team_permissions(): assert exc_info.value.type == ProxyErrorTypes.team_vector_store_access_denied +@pytest.mark.asyncio +@pytest.mark.parametrize( + "requested_vector_store_id,expected_error_type", + [ + ("KBOTHERTEAM99", ProxyErrorTypes.team_vector_store_access_denied), + ("KBALLOWED123", None), + ], +) +@pytest.mark.parametrize("vector_store_registry", [VectorStoreRegistry(), None], ids=["registry", "no-registry"]) +async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query( + requested_vector_store_id: str, + expected_error_type: ProxyErrorTypes | None, + vector_store_registry: VectorStoreRegistry | None, +): + """ + /v1/rag/query carries its vector store in retrieval_config.vector_store_id, + not in tools[].vector_store_ids. The team allowlist must apply either way. + """ + request_body = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": { + "vector_store_id": requested_vector_store_id, + "custom_llm_provider": "bedrock", + }, + } + valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None) + + team_object = MagicMock() + team_object.object_permission_id = "team-permission" + + mock_prisma_client = MagicMock() + team_permissions = MagicMock() + team_permissions.vector_stores = ["KBALLOWED123"] + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", vector_store_registry), + ): + if expected_error_type is None: + result = await vector_store_access_check( + request_body=request_body, + team_object=team_object, + valid_token=valid_token, + ) + assert result is True + return + + with pytest.raises(ProxyException) as exc_info: + await vector_store_access_check( + request_body=request_body, + team_object=team_object, + valid_token=valid_token, + ) + + assert exc_info.value.type == expected_error_type + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router