From f77b3b2b5234b51aefc84d64dcb9a82fc3cc0f52 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:26:31 -0700 Subject: [PATCH] refactor(s3_vectors): embed search queries through the shared vector store executor S3 Vectors now subclasses BaseQueryEmbeddingVectorStoreConfig, so its query embedding runs through the Router executor with the request metadata instead of a private router lookup. embedding_model stays accepted as an alias of litellm_embedding_model. The router kwarg is gone from the search handler and every provider transform now that nothing but the executor fallback read it. --- .../azure_ai/vector_stores/transformation.py | 7 +- .../base_llm/vector_store/transformation.py | 13 +- .../bedrock/vector_stores/transformation.py | 2 - litellm/llms/custom_httpx/llm_http_handler.py | 8 - .../gemini/vector_stores/transformation.py | 2 - .../milvus/vector_stores/transformation.py | 7 +- .../openai/vector_stores/transformation.py | 2 - .../pg_vector/vector_stores/transformation.py | 2 - .../ragflow/vector_stores/transformation.py | 2 - .../vector_stores/transformation.py | 209 +++++-------- .../vector_stores/rag_api/transformation.py | 2 - .../search_api/transformation.py | 2 - litellm/vector_stores/main.py | 1 - .../test_s3_vectors_transformation.py | 282 +++++++++--------- tests/test_litellm/vector_stores/test_main.py | 25 +- 15 files changed, 241 insertions(+), 325 deletions(-) diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index db1a0fc89a3..7638229bc32 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -24,7 +24,6 @@ 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: @@ -121,11 +120,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: 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( vector_store_id, query_text, @@ -145,11 +143,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: 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( vector_store_id, query_text, diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index 9624a721870..a3b8bcc499c 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -153,7 +153,6 @@ class BaseVectorStoreConfig: litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: Router | None = None, ) -> tuple[str, dict]: pass @@ -166,7 +165,6 @@ class BaseVectorStoreConfig: litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: Router | None = None, ) -> tuple[str, dict]: """ Optional async version of transform_search_vector_store_request. @@ -182,7 +180,6 @@ class BaseVectorStoreConfig: litellm_logging_obj=litellm_logging_obj, litellm_params=litellm_params, extra_body=extra_body, - router=router, ) @abstractmethod @@ -271,7 +268,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: pass @@ -285,7 +281,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: return self.transform_search_vector_store_request( @@ -296,7 +291,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj=litellm_logging_obj, litellm_params=litellm_params, extra_body=extra_body, - router=router, embedding_executor=embedding_executor, ) @@ -338,11 +332,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig): query_text: str, litellm_params: Mapping[str, object], embedding_executor: VectorStoreEmbeddingExecutor | None, - router: Router | None = None, ) -> Sequence[float]: model: Final = self.query_embedding_model(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: response: Final = executor.embed(model, query_text, configuration) except Exception as e: @@ -354,11 +347,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig): query_text: str, litellm_params: Mapping[str, object], embedding_executor: VectorStoreEmbeddingExecutor | None, - router: Router | None = None, ) -> Sequence[float]: model: Final = self.query_embedding_model(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: response: Final = await executor.aembed(model, query_text, configuration) except Exception as e: @@ -408,7 +400,6 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, ) -> NoReturn: raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape") diff --git a/litellm/llms/bedrock/vector_stores/transformation.py b/litellm/llms/bedrock/vector_stores/transformation.py index bad17a2181d..2d72db0cdba 100644 --- a/litellm/llms/bedrock/vector_stores/transformation.py +++ b/litellm/llms/bedrock/vector_stores/transformation.py @@ -27,7 +27,6 @@ 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 @@ -197,7 +196,6 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict]: if isinstance(query, list): query = " ".join(query) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 0f6966b0ae2..71a598a6fe7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -184,7 +184,6 @@ 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, @@ -9709,7 +9708,6 @@ class BaseLLMHTTPHandler: timeout: float | httpx.Timeout | None = None, client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - router: "Router | None" = None, ) -> VectorStoreSearchResponse: if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig): self._pre_call_direct_vector_store_search( @@ -9760,7 +9758,6 @@ class BaseLLMHTTPHandler: litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), extra_body=extra_body, - router=router, embedding_executor=embedding_executor, ) else: @@ -9775,7 +9772,6 @@ class BaseLLMHTTPHandler: litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), extra_body=extra_body, - router=router, ) all_optional_params: Final[dict[str, object]] = dict(litellm_params) all_optional_params.update(vector_store_search_optional_params or {}) @@ -9828,7 +9824,6 @@ class BaseLLMHTTPHandler: timeout: float | httpx.Timeout | None = None, client: HTTPHandler | AsyncHTTPHandler | None = None, _is_async: bool = False, - router: "Router | None" = None, ) -> VectorStoreSearchResponse | Coroutine[object, object, VectorStoreSearchResponse]: if _is_async: return self.async_vector_store_search_handler( @@ -9844,7 +9839,6 @@ class BaseLLMHTTPHandler: extra_body=extra_body, timeout=timeout, client=client, - router=router, ) if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig): @@ -9893,7 +9887,6 @@ class BaseLLMHTTPHandler: litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), extra_body=extra_body, - router=router, embedding_executor=embedding_executor, ) else: @@ -9908,7 +9901,6 @@ class BaseLLMHTTPHandler: litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), extra_body=extra_body, - router=router, ) all_optional_params: Final[dict[str, object]] = dict(litellm_params) diff --git a/litellm/llms/gemini/vector_stores/transformation.py b/litellm/llms/gemini/vector_stores/transformation.py index 82586b1f638..f6525a449b6 100644 --- a/litellm/llms/gemini/vector_stores/transformation.py +++ b/litellm/llms/gemini/vector_stores/transformation.py @@ -33,7 +33,6 @@ 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 @@ -169,7 +168,6 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: Mapping[str, object] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict]: """ Transform search request to Gemini's generateContent format. diff --git a/litellm/llms/milvus/vector_stores/transformation.py b/litellm/llms/milvus/vector_stores/transformation.py index 4f3c366d8c1..70d3649debd 100644 --- a/litellm/llms/milvus/vector_stores/transformation.py +++ b/litellm/llms/milvus/vector_stores/transformation.py @@ -24,7 +24,6 @@ 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: @@ -129,11 +128,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: 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( vector_store_id, query_text, @@ -153,11 +151,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: Mapping[str, object], extra_body: Mapping[str, object] | None = None, - router: Router | None = None, embedding_executor: VectorStoreEmbeddingExecutor | None = None, ) -> tuple[str, dict[str, object]]: 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( vector_store_id, query_text, diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 4e925494039..f6c093f2e2a 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -21,7 +21,6 @@ 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: @@ -100,7 +99,6 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict]: 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" diff --git a/litellm/llms/pg_vector/vector_stores/transformation.py b/litellm/llms/pg_vector/vector_stores/transformation.py index 9de1f589ae4..e4b06c36bf4 100644 --- a/litellm/llms/pg_vector/vector_stores/transformation.py +++ b/litellm/llms/pg_vector/vector_stores/transformation.py @@ -8,7 +8,6 @@ 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 @@ -81,7 +80,6 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict]: 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" diff --git a/litellm/llms/ragflow/vector_stores/transformation.py b/litellm/llms/ragflow/vector_stores/transformation.py index ffa6c9e1076..282cb7a92a7 100644 --- a/litellm/llms/ragflow/vector_stores/transformation.py +++ b/litellm/llms/ragflow/vector_stores/transformation.py @@ -17,7 +17,6 @@ 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 @@ -93,7 +92,6 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: dict[str, Any] | None = None, - router: "Router | None" = 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") diff --git a/litellm/llms/s3_vectors/vector_stores/transformation.py b/litellm/llms/s3_vectors/vector_stores/transformation.py index 733358381fe..a9902a0d27c 100644 --- a/litellm/llms/s3_vectors/vector_stores/transformation.py +++ b/litellm/llms/s3_vectors/vector_stores/transformation.py @@ -1,9 +1,12 @@ +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final import httpx -from litellm.caching._embedding_router import resolve_embedding_router -from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig +from litellm.llms.base_llm.vector_store.transformation import ( + BaseQueryEmbeddingVectorStoreConfig, + VectorStoreEmbeddingExecutor, +) from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( @@ -18,16 +21,18 @@ 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 +_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.""" def __init__(self) -> None: - BaseVectorStoreConfig.__init__(self) + BaseQueryEmbeddingVectorStoreConfig.__init__(self) BaseAWSLLM.__init__(self) def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials: @@ -59,141 +64,94 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): return headers 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")) return f"https://s3vectors.{aws_region_name}.api.aws" - def _resolve_query_embedding_router(self, embedding_model: str, router: "Router | None") -> "Router | None": - """Return the router iff it serves ``embedding_model`` as a deployment.""" - if router is None: - return None - model_list: Final = [ - dict(m) for m in (router.get_model_list() or ()) - ] # mutable-ok: resolve_embedding_router requires list[dict] - return resolve_embedding_router(embedding_model=embedding_model, llm_router=router, llm_model_list=model_list) + @staticmethod + def query_embedding_model(litellm_params: Mapping[str, object]) -> str: + configured: Final = litellm_params.get("litellm_embedding_model") or litellm_params.get("embedding_model") + return configured if isinstance(configured, str) and configured else _DEFAULT_QUERY_EMBEDDING_MODEL + + @staticmethod + def _query_target(vector_store_id: str, litellm_params: Mapping[str, object]) -> tuple[str, str]: + 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( self, vector_store_id: str, - query: str | list[str], + query: str | Sequence[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, - litellm_params: dict, - extra_body: dict[str, Any] | None = None, - router: "Router | None" = None, - ) -> tuple[str, dict]: - """Sync version - generates embedding synchronously.""" - # For S3 Vectors, vector_store_id should be in format: bucket_name:index_name - # If not in that format, try to construct it from litellm_params - bucket_name: str - index_name: str - - if ":" in vector_store_id: - bucket_name, index_name = vector_store_id.split(":", 1) - else: - # Try to get bucket_name from litellm_params - bucket_name_from_params: Final = litellm_params.get("vector_bucket_name") - 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) + litellm_params: Mapping[str, object], + extra_body: Mapping[str, object] | None = None, + embedding_executor: VectorStoreEmbeddingExecutor | None = None, + ) -> tuple[str, dict[str, object]]: + bucket_name, index_name = self._query_target(vector_store_id, litellm_params) + query_text: Final = self.query_text(query) + query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor) + return self._query_request( + bucket_name, + index_name, + query_text, + query_vector, + vector_store_search_optional_params, + api_base, + litellm_logging_obj, ) - 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( self, vector_store_id: str, - query: str | list[str], + query: str | Sequence[str], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, api_base: str, litellm_logging_obj: LiteLLMLoggingObj, - litellm_params: dict, - extra_body: dict[str, Any] | None = None, - router: "Router | None" = None, - ) -> tuple[str, dict]: - """Async version - generates embedding asynchronously.""" - # For S3 Vectors, vector_store_id should be in format: bucket_name:index_name - # If not in that format, try to construct it from litellm_params - bucket_name: str - index_name: str - - if ":" in vector_store_id: - bucket_name, index_name = vector_store_id.split(":", 1) - else: - # Try to get bucket_name from litellm_params - bucket_name_from_params: Final = litellm_params.get("vector_bucket_name") - 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) + litellm_params: Mapping[str, object], + extra_body: Mapping[str, object] | None = None, + embedding_executor: VectorStoreEmbeddingExecutor | None = None, + ) -> tuple[str, dict[str, object]]: + bucket_name, index_name = self._query_target(vector_store_id, litellm_params) + query_text: Final = self.query_text(query) + query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor) + return self._query_request( + bucket_name, + index_name, + query_text, + query_vector, + vector_store_search_optional_params, + api_base, + litellm_logging_obj, ) - 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( self, @@ -226,21 +184,13 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): if not source_text: continue - # Extract file information from metadata chunk_index = metadata.get("chunk_index", "0") file_id = f"s3-vectors-chunk-{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") score = 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))) results.append( @@ -265,7 +215,6 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): headers=response.headers, ) - # Vector store creation is not yet implemented def transform_create_vector_store_request( self, vector_store_create_optional_params, diff --git a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py index 36b57e7c995..5c250fc1a7e 100644 --- a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py @@ -21,7 +21,6 @@ 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: @@ -162,7 +161,6 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: Mapping[str, object] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict[str, object]]: """ Transform search request for Vertex AI RAG API diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index f0812e3ed9f..0bcf16ee06f 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -25,7 +25,6 @@ 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: @@ -246,7 +245,6 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, extra_body: Mapping[str, object] | None = None, - router: "Router | None" = None, ) -> tuple[str, dict[str, object]]: """ Transform a search request for the Vertex AI Search (Discovery Engine) API. diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index 636bdd4b52e..2fe1965a192 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -482,7 +482,6 @@ def search( timeout=timeout or request_timeout, _is_async=_is_async, client=kwargs.get("client"), - router=router, ) return response diff --git a/tests/test_litellm/llms/s3_vectors/vector_stores/test_s3_vectors_transformation.py b/tests/test_litellm/llms/s3_vectors/vector_stores/test_s3_vectors_transformation.py index 4b58d220623..e2b02ea2151 100644 --- a/tests/test_litellm/llms/s3_vectors/vector_stores/test_s3_vectors_transformation.py +++ b/tests/test_litellm/llms/s3_vectors/vector_stores/test_s3_vectors_transformation.py @@ -1,40 +1,72 @@ +from collections.abc import Mapping from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest +from litellm.llms.base_llm.vector_store.transformation import ( + RouterVectorStoreEmbeddingExecutor, +) from litellm.llms.s3_vectors.vector_stores.transformation import ( S3VectorsVectorStoreConfig, ) +from litellm.types.utils import EmbeddingResponse 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.""" - 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 + +def _embedding_response(vector): + return EmbeddingResponse(data=[{"embedding": vector, "index": 0, "object": "embedding"}]) + + +class _RecordingExecutor: + """Executor double recording every (model, query, configuration) it was asked to embed.""" + + def __init__(self, vector=QUERY_VECTOR): + self.vector = vector + self.calls = [] + + 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: def test_init(self): - """Test that S3VectorsVectorStoreConfig initializes correctly""" config = S3VectorsVectorStoreConfig() assert config is not None def test_get_supported_openai_params(self): - """Test that supported OpenAI params are returned""" config = S3VectorsVectorStoreConfig() params = config.get_supported_openai_params("test-model") assert "max_num_results" in params def test_get_complete_url(self): - """Test URL generation for S3 Vectors""" config = S3VectorsVectorStoreConfig() litellm_params = {"aws_region_name": "us-west-2"} url = config.get_complete_url(None, litellm_params) @@ -57,180 +89,149 @@ class TestS3VectorsVectorStoreConfig: 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!"}) def test_transform_search_request(self): - """Full request-body transformation with a router-injected embedding""" + """Full request-body transformation with the query embedded through the injected executor""" config = S3VectorsVectorStoreConfig() - mock_logging_obj = Mock() - mock_logging_obj.model_call_details = {} - router = _mock_router(["text-embedding-3-small"], sync=True) + logging_obj = _logging_obj() + executor = _RecordingExecutor() 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, + **_search_kwargs( + vector_store_search_optional_params={"max_num_results": 7}, + litellm_logging_obj=logging_obj, + embedding_executor=executor, + ) ) 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]}, + "queryVector": {"float32": QUERY_VECTOR}, "topK": 7, "returnDistance": 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 - 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).""" + async def test_atransform_search_embeds_alias_and_store_config_through_executor(self): + """The store's embedding_model alias and litellm_embedding_config reach the executor unchanged""" config = S3VectorsVectorStoreConfig() - mock_logging_obj = Mock() - mock_logging_obj.model_call_details = {} - router = _mock_router(["my-embedding-model"]) + executor = _RecordingExecutor(vector=[0.4, 0.5]) - 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 - 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, + _, request_body = await config.atransform_search_vector_store_request( + **_search_kwargs( + query=["test", "query"], + litellm_params={ + "embedding_model": "my-embedding-model", + "litellm_embedding_config": {"api_key": "store-key"}, + }, + embedding_executor=executor, ) + ) - 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): # 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 executor.calls == [("my-embedding-model", "test query", {"api_key": "store-key"})] assert request_body["queryVector"]["float32"] == [0.4, 0.5] + assert request_body["topK"] == 5 @pytest.mark.asyncio - async def test_atransform_search_without_router_uses_bare_embedding(self): - """Backward compat: no router -> bare litellm.aembedding as before""" + async def test_atransform_search_router_executor_carries_request_metadata(self): + """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() - mock_logging_obj = Mock() - mock_logging_obj.model_call_details = {} + router = MagicMock() + 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]}])) - 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={}, - extra_body=None, + _, request_body = await config.atransform_search_vector_store_request( + **_search_kwargs( + litellm_params={"embedding_model": "team-embeddings"}, + embedding_executor=RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata), ) + ) + + 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_without_executor_uses_bare_embedding(self): + """Backward compat: SDK callers without an executor keep embedding through litellm.aembedding""" + 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"]) 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""" + def test_transform_search_without_executor_uses_bare_embedding_sync(self): + """Sync twin: no executor -> bare litellm.embedding as before""" 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 - _, 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]}])) + mock_bare = MagicMock(return_value=_embedding_response([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 _, 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, + **_search_kwargs(litellm_params={"embedding_model": "my-embedding-model"}) ) - 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] def test_transform_search_request_invalid_vector_store_id(self): - """Test that invalid vector_store_id format raises error""" + """An unparseable vector_store_id raises before any embedding is generated""" config = S3VectorsVectorStoreConfig() - mock_logging_obj = Mock() - mock_logging_obj.model_call_details = {} + executor = _RecordingExecutor() with pytest.raises( ValueError, match="vector_store_id must be in format 'bucket_name:index_name'", ): config.transform_search_vector_store_request( - vector_store_id="invalid-format", - 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, + **_search_kwargs(vector_store_id="invalid-format", embedding_executor=executor) ) + 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): - """Test search response transformation""" config = S3VectorsVectorStoreConfig() mock_logging_obj = Mock() mock_logging_obj.model_call_details = {"query": "test query"} @@ -239,7 +240,7 @@ class TestS3VectorsVectorStoreConfig: mock_response.json.return_value = { "vectors": [ { - "distance": 0.05, # S3 Vectors returns distance, not score + "distance": 0.05, "metadata": { "source_text": "This is test content", "chunk_index": "0", @@ -258,23 +259,18 @@ class TestS3VectorsVectorStoreConfig: mock_response.status_code = 200 mock_response.headers = {} - result = config.transform_search_vector_store_response( - mock_response, mock_logging_obj - ) + result = config.transform_search_vector_store_response(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["search_query"] == "test query" assert len(result["data"]) == 2 - # Score should be 1 - distance (cosine similarity) - assert result["data"][0]["score"] == 0.95 # 1 - 0.05 + assert result["data"][0]["score"] == 0.95 assert result["data"][0]["content"][0]["text"] == "This is test content" 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" def test_map_openai_params(self): - """Test OpenAI parameter mapping""" config = S3VectorsVectorStoreConfig() non_default_params = {"max_num_results": 5} optional_params = {} diff --git a/tests/test_litellm/vector_stores/test_main.py b/tests/test_litellm/vector_stores/test_main.py index d01e696906a..234e0b01094 100644 --- a/tests/test_litellm/vector_stores/test_main.py +++ b/tests/test_litellm/vector_stores/test_main.py @@ -2,14 +2,17 @@ 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). +explicit named parameter that reaches the HTTP handler wrapped in the embedding +executor, 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.llms.base_llm.vector_store.transformation import ( + RouterVectorStoreEmbeddingExecutor, +) from litellm.vector_stores.main import search MOCK_SEARCH_RESPONSE = { @@ -19,17 +22,18 @@ MOCK_SEARCH_RESPONSE = { } -def test_search_threads_router_to_handler(): - """search() must pass its router param through to the HTTP handler""" +def test_search_wraps_router_into_the_handler_embedding_executor(): + """search() hands the HTTP handler a Router-backed embedding executor carrying the + request metadata, and no bare router kwarg (LIT-6750)""" mock_router = MagicMock() logger = MagicMock() 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", 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_store_search_handler", return_value=MOCK_SEARCH_RESPONSE, @@ -41,11 +45,16 @@ def test_search_threads_router_to_handler(): custom_llm_provider="s3_vectors", router=mock_router, litellm_logging_obj=logger, + litellm_metadata={"user_api_key_team_id": "team-a"}, ) assert response == MOCK_SEARCH_RESPONSE 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():