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 <mrinal@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-01 13:48:41 -07:00 • committed by GitHub
parent ec605826d4
commit a38fff6560
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 306 additions and 8 deletions

View file

@ -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,
):
"""

View file

@ -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) == ()

View file

@ -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