mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
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:
commit
66a3d24b3f
17 changed files with 270 additions and 337 deletions
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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 = {}
|
||||||
|
|
|
||||||
|
|
@ -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():
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue