mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector_stores): S3 Vectors search router bypass + rag query config drop + UI error swallow
This commit is contained in:
parent
24123269cc
commit
a1514efa21
24 changed files with 628 additions and 27 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
77
tests/test_litellm/vector_stores/test_main.py
Normal file
77
tests/test_litellm/vector_stores/test_main.py
Normal 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
|
||||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue