From c4e637cc318c2905a6f75901e5ad7f6a8cf5dd4f Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 10:33:39 -0700 Subject: [PATCH] fix(vector-stores): propagate chat retrieval authorization errors --- .../vector_store_pre_call_hook.py | 4 ++ .../test_vector_store_pre_call_hook.py | 55 ++++++++++++++++--- 2 files changed, 50 insertions(+), 9 deletions(-) diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index 13494f48cf3..390e1e4fbd5 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -173,6 +173,8 @@ class VectorStorePreCallHook(CustomLogger): }, ) except Exception as search_error: + if getattr(search_error, "status_code", None) in (401, 403): + raise verbose_logger.warning( "Vector store search failed for vector_store_id=%s, continuing without its context: %s", vector_store_id, @@ -202,6 +204,8 @@ class VectorStorePreCallHook(CustomLogger): return model, modified_messages, non_default_params except Exception as e: + if getattr(e, "status_code", None) in (401, 403): + raise verbose_logger.exception("Error in VectorStorePreCallHook: %s", e) # Return original parameters on error return model, messages, non_default_params diff --git a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index 26d769b72a2..e946f6fc4b1 100644 --- a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -1,9 +1,10 @@ import logging from collections.abc import Iterator from dataclasses import dataclass, field -from typing import Protocol +from typing import Final, Protocol import pytest +from fastapi import HTTPException import litellm from litellm._logging import verbose_logger @@ -40,9 +41,12 @@ def _search_response(text: str) -> VectorStoreSearchResponse: class RecordingRouter: failing_vector_store_ids: frozenset[str] = frozenset() calls: list[dict[str, object]] = field(default_factory=list) + search_error: Exception | None = None async def avector_store_search(self, **kwargs: object) -> VectorStoreSearchResponse: self.calls.append(kwargs) + if self.search_error is not None: + raise self.search_error vector_store_id = str(kwargs["vector_store_id"]) if vector_store_id in self.failing_vector_store_ids: raise litellm.BadRequestError( @@ -150,8 +154,11 @@ async def test_hook_searches_through_the_injected_router_with_the_request_metada @pytest.mark.asyncio -async def test_hook_skips_an_untrusted_managed_milvus_grpc_connection( +@pytest.mark.parametrize("vector_store_ids", [("legacy",), ("legacy", "safe"), ("safe", "legacy")]) +async def test_hook_rejects_an_untrusted_managed_milvus_grpc_connection( monkeypatch: pytest.MonkeyPatch, + vector_store_ids: tuple[str, ...], + warnings: list[logging.LogRecord], ) -> None: monkeypatch.setattr( litellm, @@ -173,16 +180,46 @@ async def test_hook_skips_an_untrusted_managed_milvus_grpc_connection( ], ), ) - router = RecordingRouter() + router: Final = RecordingRouter() + logging_obj: Final = FakeLoggingObj({}) - _, messages, _ = await _run_hook( - VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), - ["legacy", "safe"], - FakeLoggingObj({}), + with pytest.raises(HTTPException) as exc_info: + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + list(vector_store_ids), + logging_obj, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == ( + "This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used." ) + assert [call["vector_store_id"] for call in router.calls] == (["safe"] if vector_store_ids[0] == "safe" else []) + assert "search_results" not in logging_obj.model_call_details + assert warnings == [] - assert [call["vector_store_id"] for call in router.calls] == ["safe"] - assert messages[0]["content"] == "Context:\n\ncontext from safe\n\n" + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [401, 403]) +async def test_hook_propagates_search_authorization_errors( + registry_with: RegisterStores, + warnings: list[logging.LogRecord], + status_code: int, +) -> None: + registry_with("vs-denied", "vs-safe") + error: Final = HTTPException(status_code=status_code, detail="Vector store access denied") + router: Final = RecordingRouter(search_error=error) + + with pytest.raises(HTTPException) as exc_info: + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-denied", "vs-safe"], + FakeLoggingObj({}), + ) + + assert exc_info.value is error + assert [call["vector_store_id"] for call in router.calls] == ["vs-denied"] + assert warnings == [] @pytest.mark.asyncio