fix(vector-stores): survive a failing vector store search in the chat completions hook

One unreachable vector store used to wipe out every store's context on a
chat completion carrying vector_store_ids: the search raised, the blanket
handler returned the original messages, and the request answered with no
retrieved context at all. Each store's search now has its own handler that
warns with the vector store id and moves on to the next store.

The same loop appended every store's results to the original messages
instead of the running copy, so with two healthy stores only the last one
reached the model. It now chains through modified_messages.

The Router is injected through a ProxyRuntime protocol instead of an
in-function litellm.proxy.proxy_server import, so the hook's routing can
be driven in tests without touching proxy globals.
This commit is contained in:
mateo-berri 2026-09-02 21:55:16 -07:00
parent 7f7e0d5517
commit 3ea61c23c7
2 changed files with 269 additions and 27 deletions

View file

@ -6,7 +6,8 @@ It searches the vector store for relevant context and appends it to the messages
"""
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Final, cast
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
import litellm
import litellm.vector_stores
@ -24,10 +25,35 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.utils import PrismaClient
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
class ProxyRuntime(Protocol):
def llm_router(self) -> "Router | None": ...
def prisma_client(self) -> "PrismaClient | None": ...
@dataclass(frozen=True, slots=True)
class ProxyServerRuntime:
def llm_router(self) -> "Router | None":
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
return None
return llm_router
def prisma_client(self) -> "PrismaClient | None":
try:
from litellm.proxy.proxy_server import prisma_client
except ImportError:
return None
return prisma_client
class VectorStorePreCallHook(CustomLogger):
CONTENT_PREFIX_STRING = "Context:\n\n"
"""
@ -39,8 +65,9 @@ class VectorStorePreCallHook(CustomLogger):
3. Appends the search results as context to the messages
"""
def __init__(self):
def __init__(self, proxy_runtime: ProxyRuntime | None = None):
super().__init__()
self.proxy_runtime: Final[ProxyRuntime] = proxy_runtime or ProxyServerRuntime()
async def async_get_chat_completion_prompt(
self,
@ -79,21 +106,8 @@ class VectorStorePreCallHook(CustomLogger):
if litellm.vector_store_registry is None:
return model, messages, non_default_params
# Get prisma_client for database fallback
prisma_client = None
llm_router = None
try:
from litellm.proxy.proxy_server import (
llm_router as _llm_router,
)
from litellm.proxy.proxy_server import (
prisma_client as _prisma_client,
)
prisma_client = _prisma_client
llm_router = _llm_router
except ImportError:
pass
prisma_client: Final = self.proxy_runtime.prisma_client()
llm_router: Final = self.proxy_runtime.llm_router()
# Use database fallback to ensure synchronization across instances
vector_stores_to_run: list[
@ -136,15 +150,23 @@ class VectorStorePreCallHook(CustomLogger):
Callable[..., Awaitable[VectorStoreSearchResponse]],
litellm.vector_stores.asearch,
)
search_response = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
try:
search_response = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
except Exception as search_error:
verbose_logger.warning(
"Vector store search failed for vector_store_id=%s, continuing without its context: %s",
vector_store_id,
search_error,
)
continue
verbose_logger.debug("search_response: %s", search_response)
@ -153,7 +175,7 @@ class VectorStorePreCallHook(CustomLogger):
# Process search results and append as context
modified_messages = self._append_search_results_to_messages(
messages=messages, search_response=search_response
messages=modified_messages, search_response=search_response
)
# Get the number of results for logging

View file

@ -0,0 +1,220 @@
import logging
from dataclasses import dataclass, field
from typing import Any
import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
VectorStorePreCallHook,
)
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
from litellm.vector_stores.vector_store_registry import (
LiteLLM_ManagedVectorStore,
VectorStoreRegistry,
)
def _search_response(text: str) -> VectorStoreSearchResponse:
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query="what is litellm?",
data=[
VectorStoreSearchResult(
score=1.0,
content=[VectorStoreResultContent(text=text, type="text")],
)
],
)
@dataclass
class RecordingRouter:
failing_vector_store_ids: frozenset[str] = frozenset()
calls: list[dict[str, Any]] = field(default_factory=list)
async def avector_store_search(self, **kwargs: Any) -> VectorStoreSearchResponse:
self.calls.append(kwargs)
vector_store_id = kwargs["vector_store_id"]
if vector_store_id in self.failing_vector_store_ids:
raise litellm.BadRequestError(
message=f"no healthy deployments for {vector_store_id}",
model="text-embedding-3-small",
llm_provider="openai",
)
return _search_response(f"context from {vector_store_id}")
@dataclass(frozen=True)
class FakeProxyRuntime:
router: RecordingRouter | None
def llm_router(self) -> RecordingRouter | None:
return self.router
def prisma_client(self) -> None:
return None
class RecordingHandler(logging.Handler):
def __init__(self) -> None:
super().__init__(level=logging.WARNING)
self.records: list[logging.LogRecord] = []
def emit(self, record: logging.LogRecord) -> None:
self.records.append(record)
@pytest.fixture
def registry_with(monkeypatch: pytest.MonkeyPatch):
def _register(*vector_store_ids: str, custom_llm_provider: str = "bedrock") -> None:
monkeypatch.setattr(
litellm,
"vector_store_registry",
VectorStoreRegistry(
vector_stores=[
LiteLLM_ManagedVectorStore(vector_store_id=vector_store_id, custom_llm_provider=custom_llm_provider)
for vector_store_id in vector_store_ids
],
),
)
return _register
@pytest.fixture
def warnings():
handler = RecordingHandler()
verbose_logger.addHandler(handler)
yield handler.records
verbose_logger.removeHandler(handler)
class FakeLoggingObj:
def __init__(self, metadata: dict[str, Any]) -> None:
self.model_call_details: dict[str, Any] = {"litellm_params": {"metadata": metadata}}
async def _run_hook(hook: VectorStorePreCallHook, vector_store_ids: list[str], logging_obj: FakeLoggingObj):
return await hook.async_get_chat_completion_prompt(
model="chat-model",
messages=[{"role": "user", "content": "what is litellm?"}],
non_default_params={"vector_store_ids": vector_store_ids},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
litellm_logging_obj=logging_obj,
)
@pytest.mark.asyncio
async def test_hook_searches_through_the_injected_router_with_the_request_metadata(registry_with):
"""Regression (LIT-6752): the hook must reach the Router through its injected runtime, not a proxy_server import."""
registry_with("vs-router")
router = RecordingRouter()
logging_obj = FakeLoggingObj({"user_api_key_team_id": "team-a"})
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-router"],
logging_obj,
)
assert router.calls == [
{
"vector_store_id": "vs-router",
"query": "what is litellm?",
"custom_llm_provider": "bedrock",
"metadata": {"user_api_key_team_id": "team-a"},
}
]
assert messages[0]["content"] == "Context:\n\ncontext from vs-router\n\n"
@pytest.mark.asyncio
async def test_hook_falls_back_to_the_sdk_when_the_runtime_has_no_router(registry_with, warnings):
registry_with("vs-sdk", custom_llm_provider="lit6752-not-a-provider")
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)),
["vs-sdk"],
FakeLoggingObj({"user_api_key_team_id": "team-a"}),
)
assert messages == [{"role": "user", "content": "what is litellm?"}]
assert len(warnings) == 1
assert (
warnings[0]
.getMessage()
.startswith("Vector store search failed for vector_store_id=vs-sdk, continuing without its context: ")
)
assert "is not a valid LlmProviders" in warnings[0].getMessage()
@pytest.mark.asyncio
async def test_every_healthy_vector_store_contributes_its_own_context(registry_with):
"""Regression (LIT-6752): each store appended its context to the original messages, so only the last one survived."""
registry_with("vs-one", "vs-two")
router = RecordingRouter()
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-one", "vs-two"],
FakeLoggingObj({}),
)
assert [message["content"] for message in messages] == [
"Context:\n\ncontext from vs-one\n\n",
"Context:\n\ncontext from vs-two\n\n",
"what is litellm?",
]
@pytest.mark.asyncio
async def test_a_failing_vector_store_warns_with_its_id_and_the_other_stores_still_answer(registry_with, warnings):
"""Regression (LIT-6752): one unreachable store must not silently drop every other store's context."""
registry_with("vs-broken", "vs-healthy")
router = RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))
logging_obj = FakeLoggingObj({"user_api_key_team_id": "team-a"})
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-broken", "vs-healthy"],
logging_obj,
)
assert [call["vector_store_id"] for call in router.calls] == ["vs-broken", "vs-healthy"]
assert messages[0]["content"] == "Context:\n\ncontext from vs-healthy\n\n"
assert len(logging_obj.model_call_details["search_results"]) == 1
assert [record.getMessage() for record in warnings] == [
"Vector store search failed for vector_store_id=vs-broken, continuing without its context: "
"litellm.BadRequestError: no healthy deployments for vs-broken"
]
@pytest.mark.asyncio
async def test_the_only_vector_store_failing_leaves_the_messages_untouched(registry_with, warnings):
registry_with("vs-broken")
original_messages = [{"role": "user", "content": "what is litellm?"}]
_, messages, _ = await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken"],
FakeLoggingObj({}),
)
assert messages == original_messages
assert [(record.levelname, record.getMessage()) for record in warnings] == [
(
"WARNING",
"Vector store search failed for vector_store_id=vs-broken, continuing without its context: "
"litellm.BadRequestError: no healthy deployments for vs-broken",
)
]