Merge pull request #39474 from BerriAI/litellm_s3_vectors_query_embedding_executor

refactor(s3_vectors): embed search queries through the shared vector store executor
This commit is contained in:
Mateo Wang 2026-09-02 22:23:44 -07:00 committed by GitHub
commit 66a3d24b3f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 270 additions and 337 deletions

View file

@ -24,7 +24,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj LiteLLMLoggingObj = _LiteLLMLoggingObj
else: else:
@ -121,11 +120,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query) query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router) query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._search_request( return self._search_request(
vector_store_id, vector_store_id,
query_text, query_text,
@ -145,11 +143,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query) query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router) query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._search_request( return self._search_request(
vector_store_id, vector_store_id,
query_text, query_text,

View file

@ -99,12 +99,9 @@ class RouterVectorStoreEmbeddingExecutor:
) )
return bool(resolved) or model in deployment_models return bool(resolved) or model in deployment_models
def _embeds_through_sdk(self, model: str, configuration: Mapping[str, object]) -> bool:
return bool(configuration) and not self._router_serves(model)
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
embedding_kwargs: Final = self._embedding_kwargs(configuration) embedding_kwargs: Final = self._embedding_kwargs(configuration)
if self._embeds_through_sdk(model, configuration): if not self._router_serves(model):
return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, embedding_kwargs) return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, embedding_kwargs)
return self.router.embedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list return self.router.embedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
model=model, model=model,
@ -114,7 +111,7 @@ class RouterVectorStoreEmbeddingExecutor:
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
embedding_kwargs: Final = self._embedding_kwargs(configuration) embedding_kwargs: Final = self._embedding_kwargs(configuration)
if self._embeds_through_sdk(model, configuration): if not self._router_serves(model):
return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, embedding_kwargs) return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, embedding_kwargs)
return await self.router.aembedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list return await self.router.aembedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
model=model, model=model,
@ -153,7 +150,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: dict,
extra_body: dict[str, Any] | None = None, extra_body: dict[str, Any] | None = None,
router: Router | None = None,
) -> tuple[str, dict]: ) -> tuple[str, dict]:
pass pass
@ -166,7 +162,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: dict,
extra_body: dict[str, Any] | None = None, extra_body: dict[str, Any] | None = None,
router: Router | None = None,
) -> tuple[str, dict]: ) -> tuple[str, dict]:
""" """
Optional async version of transform_search_vector_store_request. Optional async version of transform_search_vector_store_request.
@ -182,7 +177,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj=litellm_logging_obj, litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params, litellm_params=litellm_params,
extra_body=extra_body, extra_body=extra_body,
router=router,
) )
@abstractmethod @abstractmethod
@ -271,7 +265,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
pass pass
@ -285,7 +278,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
return self.transform_search_vector_store_request( return self.transform_search_vector_store_request(
@ -296,7 +288,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj=litellm_logging_obj, litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params, litellm_params=litellm_params,
extra_body=extra_body, extra_body=extra_body,
router=router,
embedding_executor=embedding_executor, embedding_executor=embedding_executor,
) )
@ -338,11 +329,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
query_text: str, query_text: str,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None, embedding_executor: VectorStoreEmbeddingExecutor | None,
router: Router | None = None,
) -> Sequence[float]: ) -> Sequence[float]:
model: Final = self.query_embedding_model(litellm_params) model: Final = self.query_embedding_model(litellm_params)
configuration: Final = self.query_embedding_configuration(litellm_params) configuration: Final = self.query_embedding_configuration(litellm_params)
executor: Final = self.query_embedding_executor(embedding_executor, router) executor: Final = self.query_embedding_executor(embedding_executor, None)
try: try:
response: Final = executor.embed(model, query_text, configuration) response: Final = executor.embed(model, query_text, configuration)
except Exception as e: except Exception as e:
@ -354,11 +344,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
query_text: str, query_text: str,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None, embedding_executor: VectorStoreEmbeddingExecutor | None,
router: Router | None = None,
) -> Sequence[float]: ) -> Sequence[float]:
model: Final = self.query_embedding_model(litellm_params) model: Final = self.query_embedding_model(litellm_params)
configuration: Final = self.query_embedding_configuration(litellm_params) configuration: Final = self.query_embedding_configuration(litellm_params)
executor: Final = self.query_embedding_executor(embedding_executor, router) executor: Final = self.query_embedding_executor(embedding_executor, None)
try: try:
response: Final = await executor.aembed(model, query_text, configuration) response: Final = await executor.aembed(model, query_text, configuration)
except Exception as e: except Exception as e:
@ -408,7 +397,6 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
) -> NoReturn: ) -> NoReturn:
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape") raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape")

View file

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

View file

@ -184,7 +184,6 @@ if TYPE_CHECKING:
AnthropicMessagesStreamingResponse, AnthropicMessagesStreamingResponse,
) )
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.router import Router
from litellm.types.llms.openai_evals import ( from litellm.types.llms.openai_evals import (
CancelEvalResponse, CancelEvalResponse,
CancelRunResponse, CancelRunResponse,
@ -9709,7 +9708,6 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout | None = None, timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None, client: HTTPHandler | AsyncHTTPHandler | None = None,
_is_async: bool = False, _is_async: bool = False,
router: "Router | None" = None,
) -> VectorStoreSearchResponse: ) -> VectorStoreSearchResponse:
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig): if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
self._pre_call_direct_vector_store_search( self._pre_call_direct_vector_store_search(
@ -9760,7 +9758,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj, litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), litellm_params=dict(litellm_params),
extra_body=extra_body, extra_body=extra_body,
router=router,
embedding_executor=embedding_executor, embedding_executor=embedding_executor,
) )
else: else:
@ -9775,7 +9772,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj, litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), litellm_params=dict(litellm_params),
extra_body=extra_body, extra_body=extra_body,
router=router,
) )
all_optional_params: Final[dict[str, object]] = dict(litellm_params) all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {}) all_optional_params.update(vector_store_search_optional_params or {})
@ -9828,7 +9824,6 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout | None = None, timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None, client: HTTPHandler | AsyncHTTPHandler | None = None,
_is_async: bool = False, _is_async: bool = False,
router: "Router | None" = None,
) -> VectorStoreSearchResponse | Coroutine[object, object, VectorStoreSearchResponse]: ) -> VectorStoreSearchResponse | Coroutine[object, object, VectorStoreSearchResponse]:
if _is_async: if _is_async:
return self.async_vector_store_search_handler( return self.async_vector_store_search_handler(
@ -9844,7 +9839,6 @@ class BaseLLMHTTPHandler:
extra_body=extra_body, extra_body=extra_body,
timeout=timeout, timeout=timeout,
client=client, client=client,
router=router,
) )
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig): if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
@ -9893,7 +9887,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj, litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), litellm_params=dict(litellm_params),
extra_body=extra_body, extra_body=extra_body,
router=router,
embedding_executor=embedding_executor, embedding_executor=embedding_executor,
) )
else: else:
@ -9908,7 +9901,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj, litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), litellm_params=dict(litellm_params),
extra_body=extra_body, extra_body=extra_body,
router=router,
) )
all_optional_params: Final[dict[str, object]] = dict(litellm_params) all_optional_params: Final[dict[str, object]] = dict(litellm_params)

View file

@ -33,7 +33,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else: else:
LiteLLMLoggingObj = Any LiteLLMLoggingObj = Any
@ -169,7 +168,6 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: dict,
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]: ) -> tuple[str, dict]:
""" """
Transform search request to Gemini's generateContent format. Transform search request to Gemini's generateContent format.

View file

@ -24,7 +24,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj LiteLLMLoggingObj = _LiteLLMLoggingObj
else: else:
@ -129,11 +128,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query) query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router) query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._search_request( return self._search_request(
vector_store_id, vector_store_id,
query_text, query_text,
@ -153,11 +151,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object], litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query) query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router) query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._search_request( return self._search_request(
vector_store_id, vector_store_id,
query_text, query_text,

View file

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

View file

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

View file

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

View file

@ -1,9 +1,12 @@
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final from typing import TYPE_CHECKING, Any, Final
import httpx import httpx
from litellm.caching._embedding_router import resolve_embedding_router from litellm.llms.base_llm.vector_store.transformation import (
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig BaseQueryEmbeddingVectorStoreConfig,
VectorStoreEmbeddingExecutor,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.types.router import GenericLiteLLMParams from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import ( from litellm.types.vector_stores import (
@ -18,16 +21,18 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else: else:
LiteLLMLoggingObj = Any LiteLLMLoggingObj = Any
_DEFAULT_QUERY_EMBEDDING_MODEL: Final = "text-embedding-3-small"
_DEFAULT_TOP_K: Final = 5
class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM):
"""Vector store configuration for AWS S3 Vectors.""" """Vector store configuration for AWS S3 Vectors."""
def __init__(self) -> None: def __init__(self) -> None:
BaseVectorStoreConfig.__init__(self) BaseQueryEmbeddingVectorStoreConfig.__init__(self)
BaseAWSLLM.__init__(self) BaseAWSLLM.__init__(self)
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials: def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
@ -59,141 +64,94 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
return headers return headers
def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str: def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str:
# Resolve region the same way the ingestion path does:
# dynamic param -> AWS_REGION_NAME -> AWS_REGION -> default (us-west-2)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(litellm_params.get("aws_region_name")) aws_region_name: Final = 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" return f"https://s3vectors.{aws_region_name}.api.aws"
def _resolve_query_embedding_router(self, embedding_model: str, router: "Router | None") -> "Router | None": @staticmethod
"""Return the router iff it serves ``embedding_model`` as a deployment.""" def query_embedding_model(litellm_params: Mapping[str, object]) -> str:
if router is None: configured: Final = litellm_params.get("litellm_embedding_model") or litellm_params.get("embedding_model")
return None return configured if isinstance(configured, str) and configured else _DEFAULT_QUERY_EMBEDDING_MODEL
model_list: Final = [
dict(m) for m in (router.get_model_list() or ()) @staticmethod
] # mutable-ok: resolve_embedding_router requires list[dict] def _query_target(vector_store_id: str, litellm_params: Mapping[str, object]) -> tuple[str, str]:
return resolve_embedding_router(embedding_model=embedding_model, llm_router=router, llm_model_list=model_list) if ":" in vector_store_id:
bucket_name, index_name = vector_store_id.split(":", 1)
return bucket_name, index_name
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name")
if not isinstance(bucket_name_from_params, str) or not bucket_name_from_params:
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
return bucket_name_from_params, vector_store_id
@staticmethod
def _query_request(
bucket_name: str,
index_name: str,
query_text: str,
query_vector: Sequence[float],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
) -> tuple[str, dict[str, object]]:
litellm_logging_obj.model_call_details["query"] = query_text
return f"{api_base}/QueryVectors", {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": list(query_vector)},
"topK": vector_store_search_optional_params.get("max_num_results", _DEFAULT_TOP_K),
"returnDistance": True,
"returnMetadata": True,
}
def transform_search_vector_store_request( def transform_search_vector_store_request(
self, self,
vector_store_id: str, vector_store_id: str,
query: str | list[str], query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str, api_base: str,
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: Mapping[str, object],
extra_body: dict[str, Any] | None = None, extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict]: ) -> tuple[str, dict[str, object]]:
"""Sync version - generates embedding synchronously.""" bucket_name, index_name = self._query_target(vector_store_id, litellm_params)
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name query_text: Final = self.query_text(query)
# If not in that format, try to construct it from litellm_params query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
bucket_name: str return self._query_request(
index_name: str bucket_name,
index_name,
if ":" in vector_store_id: query_text,
bucket_name, index_name = vector_store_id.split(":", 1) query_vector,
else: vector_store_search_optional_params,
# Try to get bucket_name from litellm_params api_base,
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name") litellm_logging_obj,
if not bucket_name_from_params or not isinstance(bucket_name_from_params, str):
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
bucket_name = bucket_name_from_params
index_name = vector_store_id
if isinstance(query, list):
query = " ".join(query)
# Generate embedding for the query
embedding_model: Final = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router: Final = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_input: Final = [query] # mutable-ok: the embedding API takes list input
embedding_response: Final = (
embedding_router.embedding(model=embedding_model, input=embedding_input)
if embedding_router is not None
else litellm_module.embedding(model=embedding_model, input=embedding_input)
) )
query_embedding: Final = embedding_response.data[0]["embedding"]
url: Final = f"{api_base}/QueryVectors"
request_body: Final[dict[str, Any]] = {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": query_embedding},
"topK": vector_store_search_optional_params.get("max_num_results", 5), # Default to 5
"returnDistance": True,
"returnMetadata": True,
}
litellm_logging_obj.model_call_details["query"] = query
return url, request_body
async def atransform_search_vector_store_request( async def atransform_search_vector_store_request(
self, self,
vector_store_id: str, vector_store_id: str,
query: str | list[str], query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str, api_base: str,
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: Mapping[str, object],
extra_body: dict[str, Any] | None = None, extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict]: ) -> tuple[str, dict[str, object]]:
"""Async version - generates embedding asynchronously.""" bucket_name, index_name = self._query_target(vector_store_id, litellm_params)
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name query_text: Final = self.query_text(query)
# If not in that format, try to construct it from litellm_params query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
bucket_name: str return self._query_request(
index_name: str bucket_name,
index_name,
if ":" in vector_store_id: query_text,
bucket_name, index_name = vector_store_id.split(":", 1) query_vector,
else: vector_store_search_optional_params,
# Try to get bucket_name from litellm_params api_base,
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name") litellm_logging_obj,
if not bucket_name_from_params or not isinstance(bucket_name_from_params, str):
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
bucket_name = bucket_name_from_params
index_name = vector_store_id
if isinstance(query, list):
query = " ".join(query)
# Generate embedding for the query asynchronously
embedding_model: Final = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router: Final = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_input: Final = [query] # mutable-ok: the embedding API takes list input
embedding_response: Final = (
await embedding_router.aembedding(model=embedding_model, input=embedding_input)
if embedding_router is not None
else await litellm_module.aembedding(model=embedding_model, input=embedding_input)
) )
query_embedding: Final = embedding_response.data[0]["embedding"]
url: Final = f"{api_base}/QueryVectors"
request_body: Final[dict[str, Any]] = {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": query_embedding},
"topK": vector_store_search_optional_params.get("max_num_results", 5), # Default to 5
"returnDistance": True,
"returnMetadata": True,
}
litellm_logging_obj.model_call_details["query"] = query
return url, request_body
def sign_request( def sign_request(
self, self,
@ -226,21 +184,13 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
if not source_text: if not source_text:
continue continue
# Extract file information from metadata
chunk_index = metadata.get("chunk_index", "0") chunk_index = metadata.get("chunk_index", "0")
file_id = f"s3-vectors-chunk-{chunk_index}" file_id = f"s3-vectors-chunk-{chunk_index}"
filename = metadata.get("filename", f"document-{chunk_index}") filename = metadata.get("filename", f"document-{chunk_index}")
# S3 Vectors returns distance, convert to similarity score (0-1)
# Lower distance = higher similarity
# We'll normalize using 1 / (1 + distance) to get a 0-1 score
distance = item.get("distance") distance = item.get("distance")
score = None score = None
if distance is not None: if distance is not None:
# Convert distance to similarity score between 0 and 1
# For cosine distance: similarity = 1 - distance
# For euclidean: use 1 / (1 + distance)
# Assuming cosine distance here
score = max(0.0, min(1.0, 1.0 - float(distance))) score = max(0.0, min(1.0, 1.0 - float(distance)))
results.append( results.append(
@ -265,7 +215,6 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
headers=response.headers, headers=response.headers,
) )
# Vector store creation is not yet implemented
def transform_create_vector_store_request( def transform_create_vector_store_request(
self, self,
vector_store_create_optional_params, vector_store_create_optional_params,

View file

@ -21,7 +21,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj LiteLLMLoggingObj = _LiteLLMLoggingObj
else: else:
@ -162,7 +161,6 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: dict,
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
""" """
Transform search request for Vertex AI RAG API Transform search request for Vertex AI RAG API

View file

@ -25,7 +25,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj LiteLLMLoggingObj = _LiteLLMLoggingObj
else: else:
@ -246,7 +245,6 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj, litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict, litellm_params: dict,
extra_body: Mapping[str, object] | None = None, extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict[str, object]]: ) -> tuple[str, dict[str, object]]:
""" """
Transform a search request for the Vertex AI Search (Discovery Engine) API. Transform a search request for the Vertex AI Search (Discovery Engine) API.

View file

@ -482,7 +482,6 @@ def search(
timeout=timeout or request_timeout, timeout=timeout or request_timeout,
_is_async=_is_async, _is_async=_is_async,
client=kwargs.get("client"), client=kwargs.get("client"),
router=router,
) )
return response return response

View file

@ -376,7 +376,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
timeout=None, timeout=None,
client=None, client=None,
_is_async=False, _is_async=False,
router: "litellm.Router | None" = None, embedding_executor=None,
): ):
litellm_params_dict = ( litellm_params_dict = (
litellm_params.model_dump(exclude_none=False) litellm_params.model_dump(exclude_none=False)

View file

@ -187,7 +187,7 @@ class TestRouterEmbeddingIntegration:
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-large", ["async query"]) assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-large", ["async query"])
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_router_executor_rejects_unserved_models_without_explicit_config( async def test_router_executor_embeds_unserved_models_through_the_sdk(
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
): ):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
@ -198,12 +198,13 @@ class TestRouterEmbeddingIntegration:
metadata={"user_api_key_team_id": "team-a"}, metadata={"user_api_key_team_id": "team-a"},
) )
with pytest.raises(litellm.BadRequestError): sync_response = executor.embed("text-embedding-3-large", "sync query", {})
executor.embed("openai/text-embedding-3-large", "sync query", {}) async_response = await executor.aembed("text-embedding-3-large", "async query", {})
with pytest.raises(litellm.BadRequestError):
await executor.aembed("openai/text-embedding-3-large", "async query", {})
assert openai_route.call_count == 0 assert sync_response.data[0]["embedding"] == QUERY_VECTOR
assert async_response.data[0]["embedding"] == QUERY_VECTOR
assert _sent(openai_route, 0) == ("Bearer env-key", "text-embedding-3-large", ["sync query"])
assert _sent(openai_route, 1) == ("Bearer env-key", "text-embedding-3-large", ["async query"])
def test_router_executor_routes_deployment_model_names_through_the_router( def test_router_executor_routes_deployment_model_names_through_the_router(
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch

View file

@ -1,40 +1,70 @@
from collections.abc import Mapping
from unittest.mock import AsyncMock, MagicMock, Mock, patch from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx import httpx
import pytest import pytest
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
)
from litellm.llms.s3_vectors.vector_stores.transformation import ( from litellm.llms.s3_vectors.vector_stores.transformation import (
S3VectorsVectorStoreConfig, S3VectorsVectorStoreConfig,
) )
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import VectorStoreSearchResponse from litellm.types.vector_stores import VectorStoreSearchResponse
QUERY_VECTOR = [0.1, 0.2, 0.3]
def _mock_router(model_names, sync=False):
"""Router mock serving the given embedding model names.""" def _embedding_response(vector):
router = MagicMock() return EmbeddingResponse(data=[{"embedding": vector, "index": 0, "object": "embedding"}])
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: class _RecordingExecutor:
router.embedding = MagicMock(return_value=embedding_response) def __init__(self, vector=QUERY_VECTOR):
else: self.vector = vector
router.aembedding = AsyncMock(return_value=embedding_response) self.calls = []
return router
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
self.calls.append((model, query, dict(configuration)))
return _embedding_response(self.vector)
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
self.calls.append((model, query, dict(configuration)))
return _embedding_response(self.vector)
def _logging_obj():
logging_obj = Mock()
logging_obj.model_call_details = {}
return logging_obj
def _search_kwargs(**overrides):
kwargs = {
"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": _logging_obj(),
"litellm_params": {},
"extra_body": None,
}
kwargs.update(overrides)
return kwargs
class TestS3VectorsVectorStoreConfig: class TestS3VectorsVectorStoreConfig:
def test_init(self): def test_init(self):
"""Test that S3VectorsVectorStoreConfig initializes correctly"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
assert config is not None assert config is not None
def test_get_supported_openai_params(self): def test_get_supported_openai_params(self):
"""Test that supported OpenAI params are returned"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
params = config.get_supported_openai_params("test-model") params = config.get_supported_openai_params("test-model")
assert "max_num_results" in params assert "max_num_results" in params
def test_get_complete_url(self): def test_get_complete_url(self):
"""Test URL generation for S3 Vectors"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
litellm_params = {"aws_region_name": "us-west-2"} litellm_params = {"aws_region_name": "us-west-2"}
url = config.get_complete_url(None, litellm_params) url = config.get_complete_url(None, litellm_params)
@ -57,180 +87,170 @@ class TestS3VectorsVectorStoreConfig:
assert url == "https://s3vectors.eu-west-1.api.aws" assert url == "https://s3vectors.eu-west-1.api.aws"
def test_get_complete_url_invalid_region_format(self): def test_get_complete_url_invalid_region_format(self):
"""Invalid region format raises"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
with pytest.raises(ValueError, match="Invalid AWS region format"): with pytest.raises(ValueError, match="Invalid AWS region format"):
config.get_complete_url(None, {"aws_region_name": "Bad_Region!"}) config.get_complete_url(None, {"aws_region_name": "Bad_Region!"})
def test_transform_search_request(self): def test_transform_search_request(self):
"""Full request-body transformation with a router-injected embedding"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock() logging_obj = _logging_obj()
mock_logging_obj.model_call_details = {} executor = _RecordingExecutor()
router = _mock_router(["text-embedding-3-small"], sync=True)
url, request_body = config.transform_search_vector_store_request( url, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index", **_search_kwargs(
query="test query", vector_store_search_optional_params={"max_num_results": 7},
vector_store_search_optional_params={"max_num_results": 7}, litellm_logging_obj=logging_obj,
api_base="https://s3vectors.us-west-2.api.aws", embedding_executor=executor,
litellm_logging_obj=mock_logging_obj, )
litellm_params={},
extra_body=None,
router=router,
) )
assert url == "https://s3vectors.us-west-2.api.aws/QueryVectors" assert url == "https://s3vectors.us-west-2.api.aws/QueryVectors"
assert request_body == { assert request_body == {
"vectorBucketName": "test-bucket", "vectorBucketName": "test-bucket",
"indexName": "test-index", "indexName": "test-index",
"queryVector": {"float32": [0.1, 0.2, 0.3]}, "queryVector": {"float32": QUERY_VECTOR},
"topK": 7, "topK": 7,
"returnDistance": True, "returnDistance": True,
"returnMetadata": True, "returnMetadata": True,
} }
assert mock_logging_obj.model_call_details["query"] == "test query" assert executor.calls == [("text-embedding-3-small", "test query", {})]
assert logging_obj.model_call_details["query"] == "test query"
@pytest.mark.parametrize(
("litellm_params", "expected_model"),
[
({}, "text-embedding-3-small"),
({"embedding_model": ""}, "text-embedding-3-small"),
({"embedding_model": "my-embedding-model"}, "my-embedding-model"),
({"litellm_embedding_model": "shared-key-model"}, "shared-key-model"),
(
{"litellm_embedding_model": "shared-key-model", "embedding_model": "legacy-alias"},
"shared-key-model",
),
],
)
def test_query_embedding_model_accepts_embedding_model_alias(self, litellm_params, expected_model):
assert S3VectorsVectorStoreConfig.query_embedding_model(litellm_params) == expected_model
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_atransform_search_uses_router_for_virtual_model(self): async def test_atransform_search_embeds_alias_and_store_config_through_executor(self):
"""Regression: router-served embedding models must resolve via the router,
not a bare litellm.aembedding call (which has no deployment credentials)."""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock() executor = _RecordingExecutor(vector=[0.4, 0.5])
mock_logging_obj.model_call_details = {}
router = _mock_router(["my-embedding-model"])
with patch("litellm.aembedding", new=AsyncMock()) as mock_bare_aembedding: # test-quality-ok: guards that the bare-embedding path is not taken; dispatch seam is the behavior under test _, request_body = await config.atransform_search_vector_store_request(
url, request_body = await config.atransform_search_vector_store_request( **_search_kwargs(
vector_store_id="test-bucket:test-index", query=["test", "query"],
query="test query", litellm_params={
vector_store_search_optional_params={}, "embedding_model": "my-embedding-model",
api_base="https://s3vectors.us-west-2.api.aws", "litellm_embedding_config": {"api_key": "store-key"},
litellm_logging_obj=mock_logging_obj, },
litellm_params={"embedding_model": "my-embedding-model"}, embedding_executor=executor,
extra_body=None,
router=router,
) )
)
router.aembedding.assert_awaited_once_with(model="my-embedding-model", input=["test query"]) assert executor.calls == [("my-embedding-model", "test query", {"api_key": "store-key"})]
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): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, 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] assert request_body["queryVector"]["float32"] == [0.4, 0.5]
assert request_body["topK"] == 5
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_atransform_search_without_router_uses_bare_embedding(self): async def test_atransform_search_router_executor_carries_request_metadata(self):
"""Backward compat: no router -> bare litellm.aembedding as before""" """Regression (LIT-6750): a bare Router alias resolves through the Router with the request's
team metadata on the embedding call, so the embedding is attributed to the calling key and team."""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock() router = MagicMock()
mock_logging_obj.model_call_details = {} router.aembedding = AsyncMock(return_value=_embedding_response(QUERY_VECTOR))
request_metadata = {"user_api_key_team_id": "team-a", "user_api_key": "hashed-key"}
mock_bare = AsyncMock(return_value=Mock(data=[{"embedding": [0.6, 0.7]}])) _, request_body = await config.atransform_search_vector_store_request(
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on **_search_kwargs(
_, request_body = await config.atransform_search_vector_store_request( litellm_params={"embedding_model": "team-embeddings"},
vector_store_id="test-bucket:test-index", embedding_executor=RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata),
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,
) )
)
router.aembedding.assert_awaited_once_with(
model="team-embeddings", input=["test query"], metadata=request_metadata
)
assert request_body["queryVector"]["float32"] == QUERY_VECTOR
@pytest.mark.asyncio
async def test_atransform_search_default_model_falls_back_to_the_sdk(self):
"""Regression (LIT-6750): a store that never named an embedding model keeps working on a proxy
whose model list has no text-embedding-3-small, embedding through the SDK instead of erroring."""
config = S3VectorsVectorStoreConfig()
router = MagicMock()
router.get_model_list.return_value = [
{"model_name": "team-embeddings", "litellm_params": {"model": "openai/text-embedding-3-small"}}
]
router.resolved_litellm_models.return_value = []
router.aembedding = AsyncMock(side_effect=AssertionError("unserved model must not reach the Router"))
request_metadata = {"user_api_key_team_id": "team-a"}
mock_bare = AsyncMock(return_value=_embedding_response(QUERY_VECTOR))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose call the test asserts on
_, request_body = await config.atransform_search_vector_store_request(
**_search_kwargs(
embedding_executor=RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata)
)
)
mock_bare.assert_awaited_once_with(
model="text-embedding-3-small", input=["test query"], metadata=request_metadata
)
assert request_body["queryVector"]["float32"] == QUERY_VECTOR
@pytest.mark.asyncio
async def test_atransform_search_without_executor_uses_bare_embedding(self):
config = S3VectorsVectorStoreConfig()
mock_bare = AsyncMock(return_value=_embedding_response([0.6, 0.7]))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = await config.atransform_search_vector_store_request(**_search_kwargs())
mock_bare.assert_awaited_once_with(model="text-embedding-3-small", input=["test query"]) mock_bare.assert_awaited_once_with(model="text-embedding-3-small", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.6, 0.7] assert request_body["queryVector"]["float32"] == [0.6, 0.7]
def test_transform_search_uses_router_for_virtual_model_sync(self): def test_transform_search_without_executor_uses_bare_embedding_sync(self):
"""Sync twin: router-served embedding model resolves via router.embedding"""
config = S3VectorsVectorStoreConfig() 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: # test-quality-ok: guards that the bare-embedding path is not taken; dispatch seam is the behavior under test mock_bare = MagicMock(return_value=_embedding_response([0.8, 0.9]))
_, 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): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on with patch("litellm.embedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = config.transform_search_vector_store_request( _, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index", **_search_kwargs(litellm_params={"embedding_model": "my-embedding-model"})
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"]) mock_bare.assert_called_once_with(model="my-embedding-model", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.8, 0.9] assert request_body["queryVector"]["float32"] == [0.8, 0.9]
def test_transform_search_request_invalid_vector_store_id(self): def test_transform_search_request_invalid_vector_store_id(self):
"""Test that invalid vector_store_id format raises error"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock() executor = _RecordingExecutor()
mock_logging_obj.model_call_details = {}
with pytest.raises( with pytest.raises(
ValueError, ValueError,
match="vector_store_id must be in format 'bucket_name:index_name'", match="vector_store_id must be in format 'bucket_name:index_name'",
): ):
config.transform_search_vector_store_request( config.transform_search_vector_store_request(
vector_store_id="invalid-format", **_search_kwargs(vector_store_id="invalid-format", embedding_executor=executor)
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,
) )
assert executor.calls == []
def test_transform_search_request_bucket_from_litellm_params(self):
config = S3VectorsVectorStoreConfig()
_, request_body = config.transform_search_vector_store_request(
**_search_kwargs(
vector_store_id="only-index",
litellm_params={"vector_bucket_name": "params-bucket"},
embedding_executor=_RecordingExecutor(),
)
)
assert request_body["vectorBucketName"] == "params-bucket"
assert request_body["indexName"] == "only-index"
def test_transform_search_response(self): def test_transform_search_response(self):
"""Test search response transformation"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock() mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {"query": "test query"} mock_logging_obj.model_call_details = {"query": "test query"}
@ -239,7 +259,7 @@ class TestS3VectorsVectorStoreConfig:
mock_response.json.return_value = { mock_response.json.return_value = {
"vectors": [ "vectors": [
{ {
"distance": 0.05, # S3 Vectors returns distance, not score "distance": 0.05,
"metadata": { "metadata": {
"source_text": "This is test content", "source_text": "This is test content",
"chunk_index": "0", "chunk_index": "0",
@ -258,23 +278,18 @@ class TestS3VectorsVectorStoreConfig:
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.headers = {} mock_response.headers = {}
result = config.transform_search_vector_store_response( result = config.transform_search_vector_store_response(mock_response, mock_logging_obj)
mock_response, mock_logging_obj
)
# VectorStoreSearchResponse is a TypedDict, so check structure instead of isinstance
assert result["object"] == "vector_store.search_results.page" assert result["object"] == "vector_store.search_results.page"
assert result["search_query"] == "test query" assert result["search_query"] == "test query"
assert len(result["data"]) == 2 assert len(result["data"]) == 2
# Score should be 1 - distance (cosine similarity) assert result["data"][0]["score"] == 0.95
assert result["data"][0]["score"] == 0.95 # 1 - 0.05
assert result["data"][0]["content"][0]["text"] == "This is test content" assert result["data"][0]["content"][0]["text"] == "This is test content"
assert result["data"][0]["filename"] == "test.pdf" assert result["data"][0]["filename"] == "test.pdf"
assert result["data"][1]["score"] == 0.85 # 1 - 0.15 assert result["data"][1]["score"] == 0.85
assert result["data"][1]["content"][0]["text"] == "More test content" assert result["data"][1]["content"][0]["text"] == "More test content"
def test_map_openai_params(self): def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
config = S3VectorsVectorStoreConfig() config = S3VectorsVectorStoreConfig()
non_default_params = {"max_num_results": 5} non_default_params = {"max_num_results": 5}
optional_params = {} optional_params = {}

View file

@ -2,14 +2,17 @@
Tests for litellm/vector_stores/main.py. Tests for litellm/vector_stores/main.py.
Pins the router threading contract for vector store search: the router is an 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 explicit named parameter that reaches the HTTP handler wrapped in the embedding
into litellm_params/kwargs where logging would model_dump() it (the #19550 executor, and it must never leak into litellm_params/kwargs where logging would
serialization trap). model_dump() it (the #19550 serialization trap).
""" """
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import litellm.vector_stores.main as vector_stores_main import litellm.vector_stores.main as vector_stores_main
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
)
from litellm.vector_stores.main import search from litellm.vector_stores.main import search
MOCK_SEARCH_RESPONSE = { MOCK_SEARCH_RESPONSE = {
@ -19,17 +22,18 @@ MOCK_SEARCH_RESPONSE = {
} }
def test_search_threads_router_to_handler(): def test_search_wraps_router_into_the_handler_embedding_executor():
"""search() must pass its router param through to the HTTP handler""" """search() hands the HTTP handler a Router-backed embedding executor carrying the
request metadata, and no bare router kwarg (LIT-6750)"""
mock_router = MagicMock() mock_router = MagicMock()
logger = MagicMock() logger = MagicMock()
with ( with (
patch( # test-quality-ok: stubs provider config resolution; the seam under test is the router kwarg threading patch( # test-quality-ok: stubs provider config resolution; the seam under test is the executor threading
"litellm.vector_stores.main.ProviderConfigManager.get_provider_vector_stores_config", "litellm.vector_stores.main.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(), return_value=MagicMock(),
), ),
patch.object( # test-quality-ok: the handler call is the observable boundary for the router kwarg contract patch.object( # test-quality-ok: the handler call is the observable boundary for the executor contract
vector_stores_main.base_llm_http_handler, vector_stores_main.base_llm_http_handler,
"vector_store_search_handler", "vector_store_search_handler",
return_value=MOCK_SEARCH_RESPONSE, return_value=MOCK_SEARCH_RESPONSE,
@ -41,11 +45,16 @@ def test_search_threads_router_to_handler():
custom_llm_provider="s3_vectors", custom_llm_provider="s3_vectors",
router=mock_router, router=mock_router,
litellm_logging_obj=logger, litellm_logging_obj=logger,
litellm_metadata={"user_api_key_team_id": "team-a"},
) )
assert response == MOCK_SEARCH_RESPONSE assert response == MOCK_SEARCH_RESPONSE
mock_handler.assert_called_once() mock_handler.assert_called_once()
assert mock_handler.call_args.kwargs["router"] is mock_router assert "router" not in mock_handler.call_args.kwargs
executor = mock_handler.call_args.kwargs["embedding_executor"]
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
assert executor.router is mock_router
assert dict(executor.metadata) == {"user_api_key_team_id": "team-a"}
def test_search_router_not_in_litellm_params(): def test_search_router_not_in_litellm_params():