mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
ec605826d4
commit
a38fff6560
3 changed files with 306 additions and 8 deletions
|
|
@ -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,
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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) == ()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue