fix(vector_stores): S3 Vectors search router bypass + rag query config drop + UI error swallow

This commit is contained in:
michelligabriele 2026-07-27 17:07:58 +02:00
parent 24123269cc
commit a1514efa21
No known key found for this signature in database
24 changed files with 628 additions and 27 deletions

View file

@ -19,6 +19,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -92,6 +93,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Azure AI Search API

View file

@ -16,6 +16,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
from ..chat.transformation import BaseLLMException as _BaseLLMException
@ -56,6 +57,7 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
pass
@ -68,6 +70,7 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
"""
Optional async version of transform_search_vector_store_request.
@ -83,6 +86,7 @@ class BaseVectorStoreConfig:
litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params,
extra_body=extra_body,
router=router,
)
@abstractmethod

View file

@ -27,6 +27,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -196,6 +197,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
if isinstance(query, list):
query = " ".join(query)

View file

@ -167,6 +167,7 @@ if TYPE_CHECKING:
AnthropicMessagesStreamingResponse,
)
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.router import Router
from litellm.types.llms.openai_evals import (
CancelEvalResponse,
CancelRunResponse,
@ -9409,6 +9410,7 @@ class BaseLLMHTTPHandler:
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
router: Optional["Router"] = None,
) -> VectorStoreSearchResponse:
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
@ -9443,6 +9445,7 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
)
else:
(
@ -9456,6 +9459,7 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
)
all_optional_params: Dict[str, Any] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
@ -9507,6 +9511,7 @@ class BaseLLMHTTPHandler:
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
router: Optional["Router"] = None,
) -> Union[VectorStoreSearchResponse, Coroutine[Any, Any, VectorStoreSearchResponse]]:
if _is_async:
return self.async_vector_store_search_handler(
@ -9521,6 +9526,7 @@ class BaseLLMHTTPHandler:
extra_body=extra_body,
timeout=timeout,
client=client,
router=router,
)
if client is None or not isinstance(client, HTTPHandler):
@ -9551,6 +9557,7 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
)
all_optional_params: Dict[str, Any] = dict(litellm_params)

View file

@ -31,6 +31,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -111,6 +112,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
"""
Transform search request to Gemini's generateContent format.

View file

@ -19,6 +19,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -123,6 +124,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Azure AI Search API

View file

@ -21,6 +21,7 @@ from litellm.utils import add_openai_metadata
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -99,6 +100,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id")
url = f"{api_base}/{encoded_vector_store_id}/search"

View file

@ -8,6 +8,7 @@ from litellm.types.vector_stores import VectorStoreSearchOptionalRequestParams
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -80,6 +81,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
encoded_vector_store_id = encode_url_path_segment(vector_store_id, field_name="vector_store_id")
url = f"{api_base}/{encoded_vector_store_id}/search"

View file

@ -17,6 +17,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -92,6 +93,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
"""RAGFlow vector stores are management-only, search is not supported."""
raise NotImplementedError("RAGFlow vector stores support dataset management only, not search/retrieval")

View file

@ -1,8 +1,8 @@
import re
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm.caching._embedding_router import resolve_embedding_router
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.types.router import GenericLiteLLMParams
@ -18,6 +18,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -58,13 +59,18 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
return headers
def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str:
aws_region_name = litellm_params.get("aws_region_name")
if not aws_region_name:
raise ValueError("aws_region_name is required for S3 Vectors")
if not re.match(r"^[a-z][a-z0-9-]*$", aws_region_name):
raise ValueError("Invalid aws_region_name format")
# Resolve region the same way the ingestion path does:
# dynamic param -> AWS_REGION_NAME -> AWS_REGION -> default (us-west-2)
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(litellm_params.get("aws_region_name"))
return f"https://s3vectors.{aws_region_name}.api.aws"
def _resolve_query_embedding_router(self, embedding_model: str, router: Optional["Router"]) -> Optional["Router"]:
"""Return the router iff it serves ``embedding_model`` as a deployment."""
if router is None:
return None
model_list = [dict(m) for m in (router.get_model_list() or [])]
return resolve_embedding_router(embedding_model=embedding_model, llm_router=router, llm_model_list=model_list)
def transform_search_vector_store_request(
self,
vector_store_id: str,
@ -74,6 +80,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
"""Sync version - generates embedding synchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
@ -99,10 +106,14 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
# Generate embedding for the query
embedding_model = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_response = litellm_module.embedding(model=embedding_model, input=[query])
if embedding_router is not None:
embedding_response = embedding_router.embedding(model=embedding_model, input=[query])
else:
embedding_response = litellm_module.embedding(model=embedding_model, input=[query])
query_embedding = embedding_response.data[0]["embedding"]
url = f"{api_base}/QueryVectors"
@ -128,6 +139,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict]:
"""Async version - generates embedding asynchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
@ -153,10 +165,14 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
# Generate embedding for the query asynchronously
embedding_model = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_response = await litellm_module.aembedding(model=embedding_model, input=[query])
if embedding_router is not None:
embedding_response = await embedding_router.aembedding(model=embedding_model, input=[query])
else:
embedding_response = await litellm_module.aembedding(model=embedding_model, input=[query])
query_embedding = embedding_response.data[0]["embedding"]
url = f"{api_base}/QueryVectors"

View file

@ -19,6 +19,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -97,6 +98,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Vertex AI RAG API

View file

@ -23,6 +23,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -197,6 +198,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
router: Optional["Router"] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform a search request for the Vertex AI Search (Discovery Engine) API.

View file

@ -26,6 +26,9 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
get_form_data,
)
from litellm.proxy.vector_store_endpoints.endpoints import (
_update_request_data_with_litellm_managed_vector_store_registry,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
@ -652,6 +655,17 @@ async def rag_query(
user_api_key_dict=user_api_key_dict,
)
# Merge litellm-managed vector store params (provider, region, embedding
# model, credentials, ...) from the registry — same source the direct
# /vector_stores/{id}/search endpoint uses. User-supplied
# retrieval_config keys win on conflict.
store_data = await _update_request_data_with_litellm_managed_vector_store_registry(
data={},
vector_store_id=retrieval_config["vector_store_id"],
user_api_key_dict=user_api_key_dict,
)
retrieval_config = {**store_data, **retrieval_config}
# Add litellm data
request_data: Dict[str, Any] = {}
request_data = await add_litellm_data_to_request(

View file

@ -59,6 +59,14 @@ INGESTION_REGISTRY: Dict[str, Type[BaseRAGIngestion]] = {
"vertex_ai": VertexAIRAGIngestion,
}
# retrieval_config keys consumed by the query pipeline itself; everything else is
# forwarded to vector_stores.asearch as provider-specific params (e.g.
# aws_region_name, embedding_model, vector_bucket_name for S3 Vectors).
# `filters`/`retrieval_filter` are reserved for the explicit filter param.
_CONSUMED_RETRIEVAL_CONFIG_KEYS = frozenset(
{"vector_store_id", "custom_llm_provider", "top_k", "filters", "retrieval_filter"}
)
def get_ingestion_class(provider: str) -> Type[BaseRAGIngestion]:
"""
@ -233,13 +241,17 @@ async def _execute_query_pipeline(
raise ValueError("No query found in messages for RAG query")
# 2. Search vector store
# Forward provider-specific retrieval_config extras (region, embedding model,
# bucket, credentials refs, ...) to the search call; kwargs win on conflict.
provider_search_params = {k: v for k, v in retrieval_config.items() if k not in _CONSUMED_RETRIEVAL_CONFIG_KEYS}
with _suppressed_sub_call_billing():
search_response = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
query=query_text,
max_num_results=retrieval_config.get("top_k", 10),
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
**kwargs,
router=router,
**{**provider_search_params, **kwargs},
)
search_provider = retrieval_config.get("custom_llm_provider", "openai")

View file

@ -5820,6 +5820,7 @@ class Router:
return await self._init_vector_store_api_endpoints(
original_function=original_function,
custom_llm_provider=custom_llm_provider,
call_type=call_type,
**kwargs,
)
elif call_type in ("afile_delete", "afile_content"):
@ -5860,6 +5861,7 @@ class Router:
self,
original_function: Callable,
custom_llm_provider: Optional[str] = None,
call_type: Optional[str] = None,
**kwargs,
):
"""
@ -5878,6 +5880,12 @@ class Router:
**kwargs,
)
# For search, pass the router so provider transforms can resolve
# router-managed embedding models (e.g. S3 Vectors query embeddings).
# Assigning into kwargs also overrides any client-supplied `router` key.
if call_type == "avector_store_search":
kwargs["router"] = self
# Otherwise, call the original function directly
return await original_function(**kwargs)

View file

@ -6,7 +6,7 @@ import asyncio
import builtins
import contextvars
from functools import partial
from typing import Any, Coroutine, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Coroutine, Dict, List, Optional, Union
import httpx
@ -28,6 +28,9 @@ from litellm.types.vector_stores import (
from litellm.utils import ProviderConfigManager, client
from litellm.vector_stores.utils import VectorStoreRequestUtils
if TYPE_CHECKING:
from litellm.router import Router
####### ENVIRONMENT VARIABLES ###################
# Initialize any necessary instances or variables here
base_llm_http_handler = BaseLLMHTTPHandler()
@ -279,6 +282,7 @@ async def asearch(
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
router: Optional["Router"] = None,
**kwargs,
) -> VectorStoreSearchResponse:
"""
@ -307,6 +311,7 @@ async def asearch(
extra_body=extra_body,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
router=router,
**kwargs,
)
@ -346,6 +351,7 @@ def search(
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
router: Optional["Router"] = None,
**kwargs,
) -> Union[VectorStoreSearchResponse, Coroutine[Any, Any, VectorStoreSearchResponse]]:
"""
@ -449,6 +455,7 @@ def search(
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
router=router,
)
return response

View file

@ -1,4 +1,4 @@
from unittest.mock import MagicMock, Mock
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import pytest
@ -9,6 +9,18 @@ from litellm.llms.s3_vectors.vector_stores.transformation import (
from litellm.types.vector_stores import VectorStoreSearchResponse
def _mock_router(model_names, sync=False):
"""Router mock serving the given embedding model names."""
router = MagicMock()
router.get_model_list.return_value = [{"model_name": name} for name in model_names]
embedding_response = Mock(data=[{"embedding": [0.1, 0.2, 0.3]}])
if sync:
router.embedding = MagicMock(return_value=embedding_response)
else:
router.aembedding = AsyncMock(return_value=embedding_response)
return router
class TestS3VectorsVectorStoreConfig:
def test_init(self):
"""Test that S3VectorsVectorStoreConfig initializes correctly"""
@ -28,19 +40,174 @@ class TestS3VectorsVectorStoreConfig:
url = config.get_complete_url(None, litellm_params)
assert url == "https://s3vectors.us-west-2.api.aws"
def test_get_complete_url_missing_region(self):
"""Test that missing region raises error"""
def test_get_complete_url_missing_region(self, monkeypatch):
"""Missing region falls back to the default region (parity with ingestion)"""
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
monkeypatch.delenv("AWS_REGION", raising=False)
config = S3VectorsVectorStoreConfig()
litellm_params = {}
with pytest.raises(ValueError, match="aws_region_name is required"):
config.get_complete_url(None, litellm_params)
url = config.get_complete_url(None, {})
assert url == "https://s3vectors.us-west-2.api.aws"
def test_get_complete_url_uses_env_region(self, monkeypatch):
"""Missing region param resolves from AWS_REGION_NAME env var"""
monkeypatch.setenv("AWS_REGION_NAME", "eu-west-1")
monkeypatch.delenv("AWS_REGION", raising=False)
config = S3VectorsVectorStoreConfig()
url = config.get_complete_url(None, {})
assert url == "https://s3vectors.eu-west-1.api.aws"
def test_get_complete_url_invalid_region_format(self):
"""Invalid region format raises"""
config = S3VectorsVectorStoreConfig()
with pytest.raises(ValueError, match="Invalid AWS region format"):
config.get_complete_url(None, {"aws_region_name": "Bad_Region!"})
@pytest.mark.skip(reason="Requires embedding API call, tested in integration tests")
def test_transform_search_request(self):
"""Test search request transformation"""
# This test requires making an actual embedding API call
# It's better tested in integration tests
pass
"""Full request-body transformation with a router-injected embedding"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["text-embedding-3-small"], sync=True)
url, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={"max_num_results": 7},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
router=router,
)
assert url == "https://s3vectors.us-west-2.api.aws/QueryVectors"
assert request_body == {
"vectorBucketName": "test-bucket",
"indexName": "test-index",
"queryVector": {"float32": [0.1, 0.2, 0.3]},
"topK": 7,
"returnDistance": True,
"returnMetadata": True,
}
assert mock_logging_obj.model_call_details["query"] == "test query"
@pytest.mark.asyncio
async def test_atransform_search_uses_router_for_virtual_model(self):
"""Regression: router-served embedding models must resolve via the router,
not a bare litellm.aembedding call (which has no deployment credentials)."""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["my-embedding-model"])
with patch("litellm.aembedding", new=AsyncMock()) as mock_bare_aembedding:
url, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "my-embedding-model"},
extra_body=None,
router=router,
)
router.aembedding.assert_awaited_once_with(model="my-embedding-model", input=["test query"])
mock_bare_aembedding.assert_not_awaited()
assert request_body["queryVector"]["float32"] == [0.1, 0.2, 0.3]
assert request_body["topK"] == 5 # default
@pytest.mark.asyncio
async def test_atransform_search_falls_back_when_router_does_not_serve_model(self):
"""Router present but embedding_model is not a router deployment ->
bare litellm.aembedding keeps working (provider-prefixed + env creds stores)."""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["some-other-model"])
mock_bare = AsyncMock(return_value=Mock(data=[{"embedding": [0.4, 0.5]}]))
with patch("litellm.aembedding", new=mock_bare):
_, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "azure/text-embedding-3-small"},
extra_body=None,
router=router,
)
mock_bare.assert_awaited_once_with(model="azure/text-embedding-3-small", input=["test query"])
router.aembedding.assert_not_awaited()
assert request_body["queryVector"]["float32"] == [0.4, 0.5]
@pytest.mark.asyncio
async def test_atransform_search_without_router_uses_bare_embedding(self):
"""Backward compat: no router -> bare litellm.aembedding as before"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
mock_bare = AsyncMock(return_value=Mock(data=[{"embedding": [0.6, 0.7]}]))
with patch("litellm.aembedding", new=mock_bare):
_, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
)
mock_bare.assert_awaited_once_with(model="text-embedding-3-small", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.6, 0.7]
def test_transform_search_uses_router_for_virtual_model_sync(self):
"""Sync twin: router-served embedding model resolves via router.embedding"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["my-embedding-model"], sync=True)
with patch("litellm.embedding", new=MagicMock()) as mock_bare_embedding:
_, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "my-embedding-model"},
extra_body=None,
router=router,
)
router.embedding.assert_called_once_with(model="my-embedding-model", input=["test query"])
mock_bare_embedding.assert_not_called()
assert request_body["queryVector"]["float32"] == [0.1, 0.2, 0.3]
def test_transform_search_without_router_uses_bare_embedding_sync(self):
"""Sync twin: no router -> bare litellm.embedding as before"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
mock_bare = MagicMock(return_value=Mock(data=[{"embedding": [0.8, 0.9]}]))
with patch("litellm.embedding", new=mock_bare):
_, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
)
mock_bare.assert_called_once_with(model="text-embedding-3-small", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.8, 0.9]
def test_transform_search_request_invalid_vector_store_id(self):
"""Test that invalid vector_store_id format raises error"""

View file

@ -327,3 +327,106 @@ def test_rag_query_stream_returns_event_stream(client_internal_user):
assert response.headers.get("content-type", "").startswith("text/event-stream")
assert '"object":"chat.completion.chunk"' in response.text
assert "data: [DONE]" in response.text
def test_rag_query_merges_managed_store_params(client_internal_user):
"""
Regression: /v1/rag/query must consult the managed vector store registry
(like the direct /v1/vector_stores/{id}/search endpoint does) so that
provider, region, embedding model, etc. don't have to be repeated in
retrieval_config. Pre-fix the registry was never read, so managed S3
Vectors stores failed with "aws_region_name is required".
"""
import litellm
from litellm.types.utils import ModelResponse
mock_vector_store = {
"vector_store_id": "s3-store",
"custom_llm_provider": "s3_vectors",
"litellm_params": {
"aws_region_name": "eu-west-1",
"embedding_model": "my-embed",
"vector_bucket_name": "bkt",
},
}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = mock_vector_store
mock_response = ModelResponse(
id="chatcmpl-test",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
model="gpt-4o-mini",
)
with patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new_callable=AsyncMock,
return_value=mock_response,
) as mock_aquery, patch.object(litellm, "vector_store_registry", mock_registry), patch(
"litellm.proxy.rag_endpoints.endpoints.assert_user_can_access_vector_store_id",
new=AsyncMock(),
), patch(
"litellm.proxy.vector_store_endpoints.endpoints.assert_user_can_access_vector_store",
new=AsyncMock(),
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"retrieval_config": {"vector_store_id": "s3-store"},
},
)
assert response.status_code == 200, response.json()
mock_aquery.assert_awaited_once()
forwarded_config = mock_aquery.await_args.kwargs["retrieval_config"]
assert forwarded_config["vector_store_id"] == "s3-store"
assert forwarded_config["custom_llm_provider"] == "s3_vectors"
assert forwarded_config["aws_region_name"] == "eu-west-1"
assert forwarded_config["embedding_model"] == "my-embed"
assert forwarded_config["vector_bucket_name"] == "bkt"
def test_rag_query_user_retrieval_config_wins_over_store(client_internal_user):
"""User-supplied retrieval_config keys must win over registry values."""
import litellm
from litellm.types.utils import ModelResponse
mock_vector_store = {
"vector_store_id": "s3-store",
"custom_llm_provider": "s3_vectors",
"litellm_params": {"aws_region_name": "eu-west-1"},
}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = mock_vector_store
mock_response = ModelResponse(
id="chatcmpl-test",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
model="gpt-4o-mini",
)
with patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new_callable=AsyncMock,
return_value=mock_response,
) as mock_aquery, patch.object(litellm, "vector_store_registry", mock_registry), patch(
"litellm.proxy.rag_endpoints.endpoints.assert_user_can_access_vector_store_id",
new=AsyncMock(),
), patch(
"litellm.proxy.vector_store_endpoints.endpoints.assert_user_can_access_vector_store",
new=AsyncMock(),
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"retrieval_config": {"vector_store_id": "s3-store", "aws_region_name": "us-east-1"},
},
)
assert response.status_code == 200, response.json()
forwarded_config = mock_aquery.await_args.kwargs["retrieval_config"]
assert forwarded_config["aws_region_name"] == "us-east-1"

View file

@ -254,6 +254,96 @@ async def test_aquery_streaming_bills_sub_call_costs_into_final_event():
assert standard_logging_object["response_cost"] >= 0.003
@pytest.mark.asyncio
async def test_aquery_forwards_provider_retrieval_config_and_router_to_search():
"""
Regression: provider-specific retrieval_config keys (aws_region_name,
embedding_model, vector_bucket_name, ...) and the router must be forwarded
to the vector store search call. Pre-fix they were silently dropped, so
/v1/rag/query failed with provider config errors (e.g. S3 Vectors
"aws_region_name is required") even when the caller supplied them.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
}
]
)
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
with patch("litellm.vector_stores.asearch", new=fake_search):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={
"vector_store_id": "bkt:idx",
"custom_llm_provider": "s3_vectors",
"top_k": 5,
"aws_region_name": "eu-west-1",
"embedding_model": "my-embed",
"vector_bucket_name": "bkt",
},
router=router,
mock_response="hi",
)
assert isinstance(response, ModelResponse)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "bkt:idx"
assert search_kwargs["custom_llm_provider"] == "s3_vectors"
assert search_kwargs["max_num_results"] == 5
assert search_kwargs["router"] is router
# provider-specific extras forwarded
assert search_kwargs["aws_region_name"] == "eu-west-1"
assert search_kwargs["embedding_model"] == "my-embed"
assert search_kwargs["vector_bucket_name"] == "bkt"
# consumed keys are not duplicated into the spread
assert "top_k" not in search_kwargs
@pytest.mark.asyncio
async def test_aquery_minimal_retrieval_config_forwards_no_extras():
"""
A minimal retrieval_config must not leak consumed keys (or invent extras)
into the vector store search call.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
with patch("litellm.vector_stores.asearch", new=fake_search):
await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi",
)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "vs_test_123"
assert search_kwargs["custom_llm_provider"] == "openai"
assert search_kwargs["router"] is None
leaked = {"top_k", "filters", "retrieval_filter", "aws_region_name", "embedding_model", "vector_bucket_name"}
assert not (leaked & set(search_kwargs.keys()))
def test_rag_call_types_are_registered():
"""
query/aquery/ingest/aingest are @client-decorated entry points, so their

View file

@ -5936,3 +5936,58 @@ async def test_acreate_batch_request_bedrock_tags_override_deployment_tags():
bedrock_tags=request_tags,
)
assert mock_sign.call_args.kwargs["data"]["tags"] == request_tags
@pytest.mark.asyncio
async def test_avector_store_search_injects_router():
"""
Regression: router.avector_store_search must pass the router down to the
SDK search call so provider transforms can resolve router-managed
embedding models (e.g. S3 Vectors query embeddings).
"""
from litellm.types.vector_stores import VectorStoreSearchResponse
mock_asearch = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
# Router.__init__ binds asearch via a local import, so patch the module
# attribute before constructing the Router.
with patch("litellm.vector_stores.main.asearch", new=mock_asearch):
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"},
}
]
)
await router.avector_store_search(
vector_store_id="v", query="q", custom_llm_provider="s3_vectors"
)
mock_asearch.assert_awaited_once()
assert mock_asearch.await_args.kwargs["router"] is router
@pytest.mark.asyncio
async def test_avector_store_create_does_not_inject_router():
"""The router injection is gated on the search call type: the create path
must keep calling the SDK without a router kwarg."""
mock_acreate = AsyncMock(return_value={"id": "vs_1", "object": "vector_store"})
# avector_store_create(model=None) resolves acreate via a local import at
# call time, so patching after Router construction works here.
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"},
}
]
)
with patch("litellm.vector_stores.main.acreate", new=mock_acreate):
await router.avector_store_create(model=None, custom_llm_provider="openai")
mock_acreate.assert_awaited_once()
assert "router" not in mock_acreate.await_args.kwargs

View file

@ -0,0 +1,77 @@
"""
Tests for litellm/vector_stores/main.py.
Pins the router threading contract for vector store search: the router is an
explicit named parameter that reaches the HTTP handler, and it must never leak
into litellm_params/kwargs where logging would model_dump() it (the #19550
serialization trap).
"""
from unittest.mock import MagicMock, patch
import litellm.vector_stores.main as vector_stores_main
from litellm.vector_stores.main import search
MOCK_SEARCH_RESPONSE = {
"object": "vector_store.search_results.page",
"search_query": "q",
"data": [],
}
def test_search_threads_router_to_handler():
"""search() must pass its router param through to the HTTP handler"""
mock_router = MagicMock()
logger = MagicMock()
with (
patch(
"litellm.vector_stores.main.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch.object(
vector_stores_main.base_llm_http_handler,
"vector_store_search_handler",
return_value=MOCK_SEARCH_RESPONSE,
) as mock_handler,
):
search(
vector_store_id="bkt:idx",
query="q",
custom_llm_provider="s3_vectors",
router=mock_router,
litellm_logging_obj=logger,
)
mock_handler.assert_called_once()
assert mock_handler.call_args.kwargs["router"] is mock_router
def test_search_router_not_in_litellm_params():
"""Regression (#19550 class): the router must stay out of GenericLiteLLMParams,
otherwise pre-call logging model_dump()s it and breaks serialization."""
mock_router = MagicMock()
logger = MagicMock()
with (
patch(
"litellm.vector_stores.main.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch.object(
vector_stores_main.base_llm_http_handler,
"vector_store_search_handler",
return_value=MOCK_SEARCH_RESPONSE,
) as mock_handler,
):
search(
vector_store_id="bkt:idx",
query="q",
custom_llm_provider="s3_vectors",
router=mock_router,
litellm_logging_obj=logger,
)
litellm_params = mock_handler.call_args.kwargs["litellm_params"]
assert "router" not in litellm_params.model_dump(exclude_none=True)
assert getattr(litellm_params, "router", None) is None

View file

@ -128,16 +128,33 @@ describe("VectorStoreTester", () => {
await waitFor(() => expect(mockSearch).toHaveBeenCalledTimes(1));
});
it("reports a failed search and keeps the history empty", async () => {
it("shows the backend error in the history when a search fails", async () => {
const user = userEvent.setup();
mockSearch.mockRejectedValue(new Error("boom"));
const errorBody = '{"error":{"message":"OpenAIException - api_key is required"}}';
mockSearch.mockRejectedValue(new Error(errorBody));
renderTester();
await user.type(queryInput(), "hello");
await user.click(searchButton());
await waitFor(() => expect(mockFromBackend).toHaveBeenCalledWith("Failed to search vector store"));
expect(screen.getByText(EMPTY_STATE)).toBeInTheDocument();
await waitFor(() => expect(mockFromBackend).toHaveBeenCalledWith(errorBody));
expect(screen.getByText(`Search failed: ${errorBody}`)).toBeInTheDocument();
expect(screen.queryByText("No results found")).not.toBeInTheDocument();
expect(screen.queryByText(EMPTY_STATE)).not.toBeInTheDocument();
// the failed query stays in the input for retry
expect(queryInput()).toHaveValue("hello");
});
it('renders "No results found" for an empty result set, not an error', async () => {
const user = userEvent.setup();
mockSearch.mockResolvedValue({ object: "vector_store.search_results.page", search_query: "hello", data: [] });
renderTester();
await user.type(queryInput(), "hello");
await user.click(searchButton());
expect(await screen.findByText("No results found")).toBeInTheDocument();
expect(screen.queryByText(/search failed/i)).not.toBeInTheDocument();
});
it("clears the search history", async () => {

View file

@ -41,6 +41,7 @@ export const VectorStoreTester: React.FC<VectorStoreTesterProps> = ({ vectorStor
{
query: string;
response: VectorStoreSearchResponse | null;
error: string | null;
timestamp: number;
}[]
>([]);
@ -60,6 +61,7 @@ export const VectorStoreTester: React.FC<VectorStoreTesterProps> = ({ vectorStor
const historyEntry = {
query,
response,
error: null,
timestamp: Date.now(),
};
@ -67,7 +69,9 @@ export const VectorStoreTester: React.FC<VectorStoreTesterProps> = ({ vectorStor
setQuery("");
} catch (error) {
console.error("Error searching vector store:", error);
NotificationsManager.fromBackend("Failed to search vector store");
const errorMessage = error instanceof Error ? error.message : String(error);
NotificationsManager.fromBackend(errorMessage);
setSearchHistory((prev) => [{ query, response: null, error: errorMessage, timestamp: Date.now() }, ...prev]);
} finally {
setIsLoading(false);
}
@ -228,6 +232,8 @@ export const VectorStoreTester: React.FC<VectorStoreTesterProps> = ({ vectorStor
);
})}
</div>
) : entry.error ? (
<div className="text-sm text-destructive break-words">Search failed: {entry.error}</div>
) : (
<div className="text-sm text-muted-foreground">No results found</div>
)}

View file

@ -6851,7 +6851,7 @@ export const vectorStoreSearchCall = async (
if (!response.ok) {
const errorData = await response.text();
await handleError(errorData);
return null;
throw new Error(errorData);
}
const data = await response.json();