From 363c0de6f70367b27925e41a2d40e91f54c2c945 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 15:13:24 -0700 Subject: [PATCH] chore(vector stores): address tenant guard followups --- litellm/proxy/_lazy_openapi_snapshot.json | 34 ++++---- .../llm_passthrough_endpoints.py | 17 ++-- litellm/proxy/rag_endpoints/endpoints.py | 34 ++++---- litellm/proxy/vector_store_endpoints/utils.py | 26 +++++-- .../test_vector_store_tenant_guard.py | 77 ++++++++++++++++--- 5 files changed, 128 insertions(+), 60 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..dfebd1baf30 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3572,7 +3572,7 @@ "/anthropic/{endpoint}": { "delete": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -3616,7 +3616,7 @@ }, "get": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -3660,7 +3660,7 @@ }, "patch": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -3704,7 +3704,7 @@ }, "post": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -3748,7 +3748,7 @@ }, "put": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "parameters": [ { "in": "path", @@ -13260,7 +13260,7 @@ "/langfuse/{endpoint}": { "delete": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -13299,7 +13299,7 @@ }, "get": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -13338,7 +13338,7 @@ }, "patch": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -13377,7 +13377,7 @@ }, "post": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -13416,7 +13416,7 @@ }, "put": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "parameters": [ { "in": "path", @@ -26883,7 +26883,7 @@ "/toolset/{toolset_name}/mcp": { "delete": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -26922,7 +26922,7 @@ }, "get": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -26961,7 +26961,7 @@ }, "head": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -27000,7 +27000,7 @@ }, "options": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -27039,7 +27039,7 @@ }, "patch": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -27078,7 +27078,7 @@ }, "post": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", @@ -27117,7 +27117,7 @@ }, "put": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "parameters": [ { "in": "path", diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ddb6717cb2e..ce103f806e1 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -48,6 +48,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( assert_user_can_access_vector_store, + get_litellm_managed_vector_store, is_allowed_to_call_vector_store_endpoint, ) from litellm.secret_managers.main import get_secret_str @@ -1927,11 +1928,11 @@ async def vertex_discovery_proxy_route( "Extracted vector store ID from endpoint: %s", vector_store_id ) - # Retrieve vector store credentials from the registry - vector_store_credentials = ( - passthrough_endpoint_router.get_vector_store_credentials( - vector_store_id=vector_store_id - ) + # Retrieve LiteLLM-managed vector store credentials if the datastore id + # is registered with LiteLLM. Unknown datastore ids keep the existing + # direct Vertex pass-through behavior. + vector_store_credentials = await get_litellm_managed_vector_store( + vector_store_id=vector_store_id ) if vector_store_credentials: @@ -1939,14 +1940,10 @@ async def vertex_discovery_proxy_route( "Found vector store credentials for ID: %s", vector_store_id ) else: - verbose_proxy_logger.warning( + verbose_proxy_logger.debug( "Vector store ID %s found in endpoint but no credentials found in registry", vector_store_id, ) - raise HTTPException( - status_code=403, - detail="Access denied: You do not have permission to access this vector store", - ) discovery_handler = get_vertex_pass_through_handler(call_type="discovery") return await _base_vertex_proxy_route( diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 9e6093a47a1..3a50b703bc4 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -31,21 +31,27 @@ router = APIRouter() def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() + payload_stack = [payload] - if isinstance(payload, dict): - for key, value in payload.items(): - if key == "vector_store_id": - if not isinstance(value, str) or not value: - raise HTTPException( - status_code=400, - detail={"error": "vector_store_id must be a non-empty string"}, - ) - vector_store_ids.add(value) - continue - vector_store_ids.update(_collect_vector_store_ids_from_payload(value)) - elif isinstance(payload, list): - for item in payload: - vector_store_ids.update(_collect_vector_store_ids_from_payload(item)) + while payload_stack: + current_payload = payload_stack.pop() + + if isinstance(current_payload, dict): + for key, value in current_payload.items(): + if key == "vector_store_id": + if not isinstance(value, str) or not value: + raise HTTPException( + status_code=400, + detail={ + "error": "vector_store_id must be a non-empty string" + }, + ) + vector_store_ids.add(value) + continue + if isinstance(value, (dict, list)): + payload_stack.append(value) + elif isinstance(current_payload, list): + payload_stack.extend(current_payload) return vector_store_ids diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 827bbba630e..c09810d06dc 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -141,7 +141,7 @@ async def get_litellm_managed_vector_store( vector_store_id: str, ) -> Optional[LiteLLM_ManagedVectorStore]: """ - Resolve a LiteLLM-managed vector store from the registry or database. + Resolve a LiteLLM-managed vector store from the registry or shared cache. Provider-native vector store IDs will not be present in either location and return None, preserving direct provider behavior while still protecting @@ -165,19 +165,31 @@ async def get_litellm_managed_vector_store( ) try: - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.auth.auth_checks import ( + get_managed_vector_store_rows_by_uuids, + ) + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if prisma_client is None: return None - row = await prisma_client.db.litellm_managedvectorstorestable.find_unique( - where={"vector_store_id": vector_store_id} + rows = await get_managed_vector_store_rows_by_uuids( + uuids=[vector_store_id], + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, ) - if row is None: + if not rows: return None - return _normalize_litellm_params(LiteLLM_ManagedVectorStore(**row.model_dump())) + return _normalize_litellm_params( + LiteLLM_ManagedVectorStore(**rows[0].model_dump()) + ) except Exception as e: verbose_proxy_logger.debug( - "Failed to resolve vector store id=%s from database: %s", + "Failed to resolve vector store id=%s from shared cache: %s", vector_store_id, e, ) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index c160c5aceb3..2d6295e4029 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException, Request, Response import litellm -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth def _mock_request() -> MagicMock: @@ -288,7 +288,53 @@ async def test_vertex_discovery_denies_other_team_vector_store_credentials(): @pytest.mark.asyncio -async def test_vertex_discovery_denies_unregistered_vector_store_id(): +async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback(): + from litellm.proxy.vector_store_endpoints.utils import ( + get_litellm_managed_vector_store, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None + cache_helper = AsyncMock( + return_value=[ + LiteLLM_ManagedVectorStoresTable( + vector_store_id="vs_cached", + custom_llm_provider="openai", + vector_store_name=None, + vector_store_description=None, + vector_store_metadata=None, + created_at=None, + updated_at=None, + litellm_credential_name=None, + litellm_params={"api_base": "https://example.com"}, + team_id="team-a", + user_id=None, + ) + ] + ) + + with ( + patch.object(litellm, "vector_store_registry", mock_registry), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + new=cache_helper, + ), + ): + vector_store = await get_litellm_managed_vector_store( + vector_store_id="vs_cached" + ) + + assert vector_store is not None + assert vector_store["vector_store_id"] == "vs_cached" + assert vector_store["team_id"] == "team-a" + cache_helper.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id(): from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( vertex_discovery_proxy_route, ) @@ -296,18 +342,25 @@ async def test_vertex_discovery_denies_unregistered_vector_store_id(): request = _mock_request() request.method = "GET" - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_vector_store_credentials", - return_value=None, + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store", + new=AsyncMock(return_value=None), + ) as mock_lookup, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route", + new=AsyncMock(return_value={"ok": True}), + ) as mock_base_route, ): - with pytest.raises(HTTPException) as exc_info: - await vertex_discovery_proxy_route( - endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", - request=request, - fastapi_response=Response(), - ) + response = await vertex_discovery_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", + request=request, + fastapi_response=Response(), + ) - assert exc_info.value.status_code == 403 + assert response == {"ok": True} + mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown") + assert mock_base_route.call_args.kwargs["router_credentials"] is None @pytest.mark.asyncio