diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 3a50b703bc4..ecb13638648 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( @@ -31,10 +32,17 @@ router = APIRouter() def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() - payload_stack = [payload] + payload_stack = [(payload, 0)] while payload_stack: - current_payload = payload_stack.pop() + current_payload, depth = payload_stack.pop() + if depth > DEFAULT_MAX_RECURSE_DEPTH: + raise HTTPException( + status_code=400, + detail={ + "error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values" + }, + ) if isinstance(current_payload, dict): for key, value in current_payload.items(): @@ -49,9 +57,9 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids.add(value) continue if isinstance(value, (dict, list)): - payload_stack.append(value) + payload_stack.append((value, depth + 1)) elif isinstance(current_payload, list): - payload_stack.extend(current_payload) + payload_stack.extend((item, depth + 1) for item in current_payload) return vector_store_ids diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 05423d9843d..86e316e7f40 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -20,29 +20,6 @@ router = APIRouter() ######################################################## -async def _check_vector_store_access( - vector_store: LiteLLM_ManagedVectorStore, - user_api_key_dict: UserAPIKeyAuth, -) -> bool: - """ - Check if the user has access to the vector store. - - Delegates to :func:`can_user_access_vector_store`, which honors: - - PROXY_ADMIN bypass - - legacy vector stores with no team_id - - key-level and team-level ``object_permission.vector_stores`` allowlists - - team_id match between key and store - """ - try: - await assert_user_can_access_vector_store( - vector_store=vector_store, - user_api_key_dict=user_api_key_dict, - ) - return True - except HTTPException: - return False - - async def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index c09810d06dc..657b520b271 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -158,11 +158,15 @@ async def get_litellm_managed_vector_store( if vector_store is not None: return _normalize_litellm_params(vector_store) except Exception as e: - verbose_proxy_logger.debug( + verbose_proxy_logger.warning( "Failed to resolve vector store id=%s from registry: %s", vector_store_id, e, ) + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e try: from litellm.proxy.auth.auth_checks import ( @@ -188,12 +192,15 @@ async def get_litellm_managed_vector_store( LiteLLM_ManagedVectorStore(**rows[0].model_dump()) ) except Exception as e: - verbose_proxy_logger.debug( + verbose_proxy_logger.warning( "Failed to resolve vector store id=%s from shared cache: %s", vector_store_id, e, ) - return None + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e async def assert_user_can_access_vector_store( diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index ae8dc602e82..346a847c5dd 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -21,6 +21,7 @@ from litellm.proxy.vector_store_endpoints.utils import ( is_allowed_to_call_vector_store_files_endpoint, ) from litellm.types.utils import LlmProviders +from litellm.types.vector_stores import LiteLLM_ManagedVectorStore if TYPE_CHECKING: from litellm.router import Router @@ -194,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, llm_router: Optional["Router"] = None, + managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None, + should_lookup_registry: bool = True, ) -> Dict: """ Update request data with model routing information from managed vector store. @@ -263,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry( return data - # Legacy path: Check vector store registry for non-managed vector stores - if litellm.vector_store_registry is not None: + # Legacy path: Check vector store registry for non-managed vector stores. + vector_store_to_run = managed_vector_store + if ( + vector_store_to_run is None + and should_lookup_registry + and litellm.vector_store_registry is not None + ): vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( vector_store_id=vector_store_id ) - if vector_store_to_run is not None: - if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get( - "custom_llm_provider" - ) - if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get( - "litellm_credential_name" - ) - if "litellm_params" in vector_store_to_run: - litellm_params = vector_store_to_run.get("litellm_params", {}) or {} - data.update(litellm_params) + + if vector_store_to_run is not None: + if "custom_llm_provider" in vector_store_to_run: + data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") + if "litellm_credential_name" in vector_store_to_run: + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) + if "litellm_params" in vector_store_to_run: + litellm_params = vector_store_to_run.get("litellm_params", {}) or {} + data.update(litellm_params) return data @@ -365,7 +372,7 @@ async def vector_store_file_create( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -379,7 +386,11 @@ async def vector_store_file_create( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -464,13 +475,17 @@ async def vector_store_file_list( data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id} data.update(query_params) data["vector_store_id"] = vector_store_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -550,7 +565,7 @@ async def vector_store_file_retrieve( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -562,7 +577,11 @@ async def vector_store_file_retrieve( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -648,7 +667,7 @@ async def vector_store_file_content( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -660,7 +679,11 @@ async def vector_store_file_content( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -746,7 +769,7 @@ async def vector_store_file_update( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id data["file_id"] = file_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -758,7 +781,11 @@ async def vector_store_file_update( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -844,7 +871,7 @@ async def vector_store_file_delete( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -856,7 +883,11 @@ async def vector_store_file_delete( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) 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 2d6295e4029..380b3963c94 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 @@ -99,7 +99,8 @@ async def test_vector_store_file_create_forces_path_id_over_body_id(): assert response == {"ok": True} assert captured_data["vector_store_id"] == "vs_path_allowed" - mock_registry.get_litellm_managed_vector_store_from_registry.assert_any_call( + assert captured_data["custom_llm_provider"] == "openai" + mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with( vector_store_id="vs_path_allowed" ) @@ -227,6 +228,25 @@ async def test_rag_ingest_denies_nested_other_team_vector_store(): mock_aingest.assert_not_called() +def test_rag_payload_scan_rejects_excessive_nesting(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1): + current["nested"] = {} + current = current["nested"] + current["vector_store_id"] = "vs_too_deep" + + with pytest.raises(HTTPException) as exc_info: + _collect_vector_store_ids_from_payload(payload) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio async def test_responses_file_search_denies_other_team_vector_store(): from litellm.proxy.common_request_processing import ( @@ -333,6 +353,24 @@ async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback cache_helper.assert_awaited_once() +@pytest.mark.asyncio +async def test_get_managed_vector_store_fails_closed_on_lookup_error(): + 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.side_effect = ( + RuntimeError("registry unavailable") + ) + + with patch.object(litellm, "vector_store_registry", mock_registry): + with pytest.raises(HTTPException) as exc_info: + await get_litellm_managed_vector_store(vector_store_id="vs_registry_only") + + assert exc_info.value.status_code == 500 + + @pytest.mark.asyncio async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id(): from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (