mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge pull request #38936 from BerriAI/litellm_fix_vector_store_request_embedding_resolution
fix(vector-store): resolve embedding credentials per request
This commit is contained in:
commit
1e6a4d98a4
17 changed files with 1336 additions and 926 deletions
|
|
@ -5,6 +5,7 @@ This hook is called before making an LLM request when a vector store is configur
|
|||
It searches the vector store for relevant context and appends it to the messages.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -80,10 +81,17 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
|
||||
# Get prisma_client for database fallback
|
||||
prisma_client = None
|
||||
llm_router = None
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
llm_router as _llm_router,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client as _prisma_client,
|
||||
)
|
||||
|
||||
prisma_client = _prisma_client
|
||||
llm_router = _llm_router
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
|
@ -114,12 +122,26 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
vector_store_id = vector_store_to_run.get("vector_store_id", "")
|
||||
custom_llm_provider = vector_store_to_run.get("custom_llm_provider")
|
||||
litellm_params_for_vector_store = vector_store_to_run.get("litellm_params", {}) or {}
|
||||
# Call litellm.vector_stores.search() with the required parameters
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
request_litellm_params = litellm_logging_obj.model_call_details.get("litellm_params", {})
|
||||
request_metadata = (
|
||||
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
|
||||
)
|
||||
if llm_router is not None:
|
||||
search_function = cast( # cast-ok: normalize router search callable
|
||||
Callable[..., Awaitable[VectorStoreSearchResponse]],
|
||||
llm_router.avector_store_search,
|
||||
)
|
||||
else:
|
||||
search_function = cast( # cast-ok: normalize SDK search callable
|
||||
Callable[..., Awaitable[VectorStoreSearchResponse]],
|
||||
litellm.vector_stores.asearch,
|
||||
)
|
||||
search_response = await search_function(
|
||||
**{
|
||||
"vector_store_id": vector_store_id,
|
||||
"query": query,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"metadata": request_metadata,
|
||||
**litellm_params_for_vector_store,
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
BaseVectorStoreAuthCredentials,
|
||||
|
|
@ -26,7 +31,7 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
||||
class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM):
|
||||
"""
|
||||
Configuration for Azure AI Search Vector Store
|
||||
|
||||
|
|
@ -110,83 +115,73 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
|||
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[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
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)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
Generates embeddings using litellm.embeddings and constructs Azure AI Search request
|
||||
"""
|
||||
# Convert query to string if it's a list
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
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)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
# Get embedding model from litellm_params (required)
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model:
|
||||
raise ValueError(
|
||||
"embedding_model is required in litellm_params for Azure AI Search. "
|
||||
"Example: litellm_params['embedding_model'] = 'azure/text-embedding-3-large'"
|
||||
)
|
||||
|
||||
embedding_config: Final = litellm_params.get("litellm_embedding_config", {})
|
||||
if not embedding_config:
|
||||
raise ValueError(
|
||||
"embedding_config is required in litellm_params for Azure AI Search. "
|
||||
"Example: litellm_params['embedding_config'] = {'api_base': 'https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/', 'api_key': 'os.environ/AZURE_API_KEY', 'api_version': '2025-09-01'}"
|
||||
)
|
||||
|
||||
# Get vector field name (defaults to contentVector)
|
||||
@staticmethod
|
||||
def _search_request(
|
||||
vector_store_id: str,
|
||||
query_text: str,
|
||||
query_vector: Sequence[float],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
vector_field: Final = litellm_params.get("azure_search_vector_field", "contentVector")
|
||||
|
||||
# Get top_k (number of results to return)
|
||||
top_k: Final = vector_store_search_optional_params.get("top_k", 10)
|
||||
|
||||
# Generate embedding for the query using litellm.embeddings
|
||||
try:
|
||||
embedding_response: Final = litellm.embedding(
|
||||
model=embedding_model,
|
||||
input=[query],
|
||||
**embedding_config,
|
||||
)
|
||||
query_vector: Final = embedding_response.data[0]["embedding"]
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
|
||||
# Azure AI Search endpoint for search
|
||||
index_name: Final = vector_store_id # vector_store_id is the index name
|
||||
url: Final = f"{api_base}/indexes/{index_name}/docs/search?api-version=2024-07-01"
|
||||
|
||||
# Build the request body for Azure AI Search with vector search
|
||||
request_body: Final = {
|
||||
"search": "*", # Get all documents (filtered by vector similarity)
|
||||
"vectorQueries": [
|
||||
{
|
||||
"vector": query_vector,
|
||||
"fields": vector_field,
|
||||
"kind": "vector",
|
||||
"k": top_k, # Number of nearest neighbors to return
|
||||
}
|
||||
],
|
||||
"select": "id,content", # Fields to return (customize based on schema)
|
||||
litellm_logging_obj.model_call_details["input"] = query_text
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = litellm_params.get("litellm_embedding_model")
|
||||
litellm_logging_obj.model_call_details["top_k"] = top_k
|
||||
return f"{api_base}/indexes/{vector_store_id}/docs/search?api-version=2024-07-01", {
|
||||
"search": "*",
|
||||
"vectorQueries": [{"vector": query_vector, "fields": vector_field, "kind": "vector", "k": top_k}],
|
||||
"select": "id,content",
|
||||
"top": top_k,
|
||||
}
|
||||
|
||||
#########################################################
|
||||
# Update logging object with details of the request
|
||||
#########################################################
|
||||
litellm_logging_obj.model_call_details["input"] = query
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = embedding_model
|
||||
litellm_logging_obj.model_call_details["top_k"] = top_k
|
||||
|
||||
return url, request_body
|
||||
|
||||
def transform_search_vector_store_response(
|
||||
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
|
||||
) -> VectorStoreSearchResponse:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,16 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from abc import abstractmethod
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, NoReturn
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
VECTOR_STORE_OPENAI_PARAMS,
|
||||
BaseVectorStoreAuthCredentials,
|
||||
|
|
@ -28,6 +34,95 @@ else:
|
|||
BaseLLMException = Any
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class VectorStoreEmbeddingExecutor(Protocol):
|
||||
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: ...
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LiteLLMVectorStoreEmbeddingExecutor:
|
||||
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
import litellm
|
||||
|
||||
return litellm.embedding( # pyright: ignore[reportCallIssue, reportUnknownMemberType, reportUnknownVariableType] # provider kwargs are intentionally dynamic
|
||||
model=model,
|
||||
input=[query], # mutable-ok: LiteLLM embedding requires a mutable input list
|
||||
**dict(configuration), # pyright: ignore[reportArgumentType] # provider-specific embedding config is validated downstream # mutable-ok: kwargs require a concrete dict
|
||||
)
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
import litellm
|
||||
|
||||
return await litellm.aembedding( # pyright: ignore[reportUnknownMemberType] # provider kwargs are intentionally dynamic
|
||||
model=model,
|
||||
input=[query], # mutable-ok: LiteLLM embedding requires a mutable input list
|
||||
**dict(configuration), # pyright: ignore[reportArgumentType] # provider-specific embedding config is validated downstream # mutable-ok: kwargs require a concrete dict
|
||||
)
|
||||
|
||||
|
||||
_REQUEST_METADATA: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def vector_store_request_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
|
||||
litellm_metadata: Final = kwargs.get("litellm_metadata")
|
||||
if isinstance(litellm_metadata, dict):
|
||||
return _REQUEST_METADATA.validate_python(litellm_metadata)
|
||||
metadata: Final = kwargs.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
return _REQUEST_METADATA.validate_python(metadata)
|
||||
return MappingProxyType({})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouterVectorStoreEmbeddingExecutor:
|
||||
router: Router
|
||||
metadata: Mapping[str, object]
|
||||
|
||||
def _embedding_kwargs(self, configuration: Mapping[str, object]) -> Mapping[str, object]:
|
||||
configured_metadata: Final = configuration.get("metadata")
|
||||
metadata: Final = {
|
||||
**(configured_metadata if isinstance(configured_metadata, Mapping) else {}),
|
||||
**self.metadata,
|
||||
}
|
||||
return {
|
||||
**{key: value for key, value in configuration.items() if key not in ("input", "metadata", "model")},
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
def _router_serves(self, model: str) -> bool:
|
||||
team_id: Final = self.metadata.get("user_api_key_team_id")
|
||||
resolved: Final = self.router.resolved_litellm_models(model, team_id if isinstance(team_id, str) else None)
|
||||
deployment_models: Final = (
|
||||
deployment.get("litellm_params", {}).get("model") for deployment in self.router.get_model_list() or ()
|
||||
)
|
||||
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:
|
||||
embedding_kwargs: Final = self._embedding_kwargs(configuration)
|
||||
if self._embeds_through_sdk(model, configuration):
|
||||
return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, embedding_kwargs)
|
||||
return self.router.embedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
|
||||
model=model,
|
||||
input=[query], # mutable-ok: Router embedding requires a mutable input list
|
||||
**embedding_kwargs, # pyright: ignore[reportArgumentType] # provider kwargs are intentionally dynamic
|
||||
)
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
embedding_kwargs: Final = self._embedding_kwargs(configuration)
|
||||
if self._embeds_through_sdk(model, configuration):
|
||||
return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, embedding_kwargs)
|
||||
return await self.router.aembedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
|
||||
model=model,
|
||||
input=[query], # mutable-ok: Router embedding requires a mutable input list
|
||||
**embedding_kwargs, # pyright: ignore[reportArgumentType] # provider kwargs are intentionally dynamic
|
||||
)
|
||||
|
||||
|
||||
class BaseVectorStoreConfig:
|
||||
def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]:
|
||||
return []
|
||||
|
|
@ -58,7 +153,7 @@ class BaseVectorStoreConfig:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
pass
|
||||
|
||||
|
|
@ -71,7 +166,7 @@ class BaseVectorStoreConfig:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Optional async version of transform_search_vector_store_request.
|
||||
|
|
@ -161,6 +256,116 @@ class BaseVectorStoreConfig:
|
|||
return 0.0, 0.0
|
||||
|
||||
|
||||
_EMPTY_EMBEDDING_CONFIGURATION: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_QUERY_VECTOR: Final = TypeAdapter(list[float])
|
||||
|
||||
|
||||
class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
|
||||
@abstractmethod
|
||||
def transform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
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
|
||||
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
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(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def query_text(query: str | Sequence[str]) -> str:
|
||||
return query if isinstance(query, str) else " ".join(query)
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_model(litellm_params: Mapping[str, object]) -> str:
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if isinstance(embedding_model, str) and embedding_model:
|
||||
return embedding_model
|
||||
raise ValueError(
|
||||
"litellm_embedding_model is required in litellm_params for this vector store. "
|
||||
"Example: litellm_params['litellm_embedding_model'] = 'openai/text-embedding-3-small'"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_configuration(litellm_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
configuration: Final = litellm_params.get("litellm_embedding_config")
|
||||
if isinstance(configuration, Mapping):
|
||||
return {str(key): value for key, value in configuration.items()} # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # litellm_params is an untyped dict, keys are re-validated as str here
|
||||
return _EMPTY_EMBEDDING_CONFIGURATION
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_executor(
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None,
|
||||
router: Router | None,
|
||||
request_metadata: Mapping[str, object] = MappingProxyType({}),
|
||||
) -> VectorStoreEmbeddingExecutor:
|
||||
if embedding_executor is not None:
|
||||
return embedding_executor
|
||||
if router is not None:
|
||||
return RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata)
|
||||
return LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
def embed_query(
|
||||
self,
|
||||
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)
|
||||
try:
|
||||
response: Final = executor.embed(model, query_text, configuration)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
return _QUERY_VECTOR.validate_python(response.data[0]["embedding"]) # pyright: ignore[reportUnknownMemberType] # EmbeddingResponse.data is an untyped list, the vector is validated here
|
||||
|
||||
async def aembed_query(
|
||||
self,
|
||||
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)
|
||||
try:
|
||||
response: Final = await executor.aembed(model, query_text, configuration)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
return _QUERY_VECTOR.validate_python(response.data[0]["embedding"]) # pyright: ignore[reportUnknownMemberType] # EmbeddingResponse.data is an untyped list, the vector is validated here
|
||||
|
||||
|
||||
class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
||||
"""
|
||||
Base config for vector store providers whose datastore has no HTTP API
|
||||
|
|
@ -176,6 +381,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
pass
|
||||
|
|
@ -188,6 +394,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
pass
|
||||
|
|
@ -201,7 +408,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> NoReturn:
|
||||
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape")
|
||||
|
||||
|
|
|
|||
|
|
@ -69,7 +69,9 @@ from litellm.llms.base_llm.skills.transformation import BaseSkillsAPIConfig
|
|||
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseDirectVectorStoreConfig,
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
BaseVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store_files.transformation import (
|
||||
BaseVectorStoreFilesConfig,
|
||||
|
|
@ -9701,6 +9703,7 @@ class BaseLLMHTTPHandler:
|
|||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
|
|
@ -9721,6 +9724,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
|
||||
embedding_executor=embedding_executor,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
|
@ -9744,8 +9748,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
# Check if provider has async transform method
|
||||
if hasattr(vector_store_provider_config, "atransform_search_vector_store_request"):
|
||||
if isinstance(vector_store_provider_config, BaseQueryEmbeddingVectorStoreConfig):
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
|
|
@ -9758,12 +9761,13 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
else:
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
) = await vector_store_provider_config.atransform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
|
|
@ -9818,6 +9822,7 @@ class BaseLLMHTTPHandler:
|
|||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
|
|
@ -9834,6 +9839,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
embedding_executor=embedding_executor,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
|
|
@ -9854,6 +9860,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
|
||||
embedding_executor=embedding_executor,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
|
@ -9874,19 +9881,35 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
)
|
||||
if isinstance(vector_store_provider_config, BaseQueryEmbeddingVectorStoreConfig):
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
else:
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
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 {})
|
||||
|
|
|
|||
|
|
@ -1,9 +1,14 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -37,7 +42,7 @@ MILVUS_OPTIONAL_PARAMS: Final = {
|
|||
}
|
||||
|
||||
|
||||
class MilvusVectorStoreConfig(BaseVectorStoreConfig):
|
||||
class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
|
||||
"""
|
||||
Configuration for Milvus Vector Store
|
||||
|
||||
|
|
@ -118,78 +123,79 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
|
|||
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[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
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)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
Generates embeddings using litellm.embeddings and constructs Azure AI Search request
|
||||
"""
|
||||
# Convert query to string if it's a list
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
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)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
# Get embedding model from litellm_params (required)
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model:
|
||||
raise ValueError(
|
||||
"embedding_model is required in litellm_params for Milvus. You can call any litellm embedding model."
|
||||
"Example: litellm_params['embedding_model'] = 'azure/text-embedding-3-large'"
|
||||
@staticmethod
|
||||
def _search_request(
|
||||
vector_store_id: str,
|
||||
query_text: str,
|
||||
query_vector: Sequence[float],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
scope: Final = {
|
||||
key: value
|
||||
for key, value in (
|
||||
("dbName", litellm_params.get("milvus_db_name")),
|
||||
("partitionNames", litellm_params.get("milvus_partition_names")),
|
||||
)
|
||||
|
||||
embedding_config: Final = litellm_params.get("litellm_embedding_config", {})
|
||||
if not embedding_config:
|
||||
raise ValueError(
|
||||
"embedding_config is required in litellm_params for Milvus. You can call any litellm embedding model."
|
||||
"Example: litellm_params['embedding_config'] = {'api_base': 'https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/', 'api_key': 'os.environ/AZURE_API_KEY', 'api_version': '2025-09-01'}"
|
||||
)
|
||||
|
||||
# Get top_k (number of results to return)
|
||||
# Generate embedding for the query using litellm.embeddings
|
||||
try:
|
||||
embedding_response: Final = litellm.embedding(
|
||||
model=embedding_model,
|
||||
input=[query],
|
||||
**embedding_config,
|
||||
)
|
||||
query_vector: Final = embedding_response.data[0]["embedding"]
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
|
||||
# Azure AI Search endpoint for search
|
||||
index_name: Final = vector_store_id # vector_store_id is the index name
|
||||
url: Final = f"{api_base}/v2/vectordb/entities/search"
|
||||
|
||||
# Build the request body for Azure AI Search with vector search
|
||||
request_body: Final[dict[str, Any]] = {
|
||||
"collectionName": index_name,
|
||||
if value
|
||||
}
|
||||
litellm_logging_obj.model_call_details["input"] = query_text
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = litellm_params.get("litellm_embedding_model")
|
||||
return f"{api_base}/v2/vectordb/entities/search", {
|
||||
"collectionName": vector_store_id,
|
||||
"data": [query_vector],
|
||||
"annsField": "book_intro_vector",
|
||||
**vector_store_search_optional_params,
|
||||
**scope,
|
||||
}
|
||||
|
||||
db_name: Final = litellm_params.get("milvus_db_name")
|
||||
if db_name:
|
||||
request_body["dbName"] = db_name
|
||||
|
||||
partition_names: Final = litellm_params.get("milvus_partition_names")
|
||||
if partition_names:
|
||||
request_body["partitionNames"] = partition_names
|
||||
|
||||
#########################################################
|
||||
# Update logging object with details of the request
|
||||
#########################################################
|
||||
litellm_logging_obj.model_call_details["input"] = query
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = embedding_model
|
||||
|
||||
return url, request_body
|
||||
|
||||
def transform_search_vector_store_response(
|
||||
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
|
||||
) -> VectorStoreSearchResponse:
|
||||
|
|
|
|||
|
|
@ -15,7 +15,10 @@ import httpx
|
|||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseDirectVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseDirectVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -213,6 +216,7 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
|
|
@ -222,10 +226,18 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = self.embedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
embedding_response: Final = (
|
||||
embedding_executor.embed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
if embedding_executor is not None
|
||||
else self.embedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: the injected embedding callable requires list input
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
|
|
@ -252,6 +264,7 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
|
|
@ -261,10 +274,18 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = await self.aembedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
embedding_response: Final = (
|
||||
await embedding_executor.aembed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
if embedding_executor is not None
|
||||
else await self.aembedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: the injected embedding callable requires list input
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
|
|
|
|||
|
|
@ -729,7 +729,7 @@ async def rag_query(
|
|||
# conflict so callers cannot override the store's provider or credentials.
|
||||
managed_store: Final = resolved_stores.get(retrieval_config["vector_store_id"])
|
||||
store_data: Final = (
|
||||
await build_request_data_from_managed_vector_store(managed_store)
|
||||
build_request_data_from_managed_vector_store(managed_store)
|
||||
if managed_store is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
|
|
|
|||
|
|
@ -16,9 +16,6 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.utils import jsonify_object
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_resolve_embedding_config,
|
||||
)
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
assert_proxy_admin_for_vector_store_index_management,
|
||||
assert_user_can_access_vector_store,
|
||||
|
|
@ -57,19 +54,9 @@ def reject_caller_embedding_selection_params(payload: Mapping[str, object], sour
|
|||
########################################################
|
||||
|
||||
|
||||
async def build_request_data_from_managed_vector_store(
|
||||
def build_request_data_from_managed_vector_store(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
Build request params (provider, credential ref, litellm_params) from an
|
||||
already-resolved managed vector store.
|
||||
|
||||
``litellm_embedding_config`` is resolved here, at request-handling time,
|
||||
instead of at row-creation time: the resolved api_key/api_base/api_version
|
||||
lives only in the returned per-request mapping and is never persisted back
|
||||
to the registry cache. Legacy rows that already carry a resolved
|
||||
(cleartext) config skip the lookup and pass through unchanged.
|
||||
"""
|
||||
top_level: Final = MappingProxyType(
|
||||
{
|
||||
key: vector_store.get(key)
|
||||
|
|
@ -78,18 +65,7 @@ async def build_request_data_from_managed_vector_store(
|
|||
}
|
||||
)
|
||||
litellm_params: Final = vector_store.get("litellm_params") or MappingProxyType({})
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model or litellm_params.get("litellm_embedding_config"):
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
resolved_config: Final = await _resolve_embedding_config(
|
||||
embedding_model=embedding_model, prisma_client=prisma_client
|
||||
)
|
||||
if not resolved_config:
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
return MappingProxyType({**top_level, **litellm_params, "litellm_embedding_config": resolved_config})
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
|
||||
|
||||
async def _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
|
|
@ -118,7 +94,7 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
vector_store=vector_store_to_run,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return {**data, **(await build_request_data_from_managed_vector_store(vector_store_to_run))}
|
||||
return {**data, **build_request_data_from_managed_vector_store(vector_store_to_run)}
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -18,11 +18,8 @@ if TYPE_CHECKING:
|
|||
from prisma.models import LiteLLM_ManagedVectorStoresTable as _VectorStoreRow
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.router import Router
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
|
@ -32,13 +29,10 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
|
||||
from litellm.repositories.model_repository import ModelRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
|
||||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.types.vector_stores import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
LiteLLM_ManagedVectorStoreListResponse,
|
||||
|
|
@ -64,28 +58,6 @@ _LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker()
|
|||
|
||||
_REDACT_LITELLM_PARAMS_MAX_DEPTH: Final = 10
|
||||
|
||||
# Use-time embedding-config resolution runs on every vector-store request
|
||||
# whose persisted row carries only a model reference (the post-fix shape).
|
||||
# Without a cache, that's one ``litellm_proxymodeltable.find_first`` per
|
||||
# request — the no-DB-in-critical-path rule. Hold the resolved config in
|
||||
# memory for a short TTL so a hot model name pays the DB lookup at most
|
||||
# once per ``_EMBEDDING_CONFIG_CACHE_TTL`` seconds. Cleartext credentials
|
||||
# only ever live in process memory (never persisted, never echoed in
|
||||
# management responses), so the cache doesn't widen the disclosure surface.
|
||||
_EMBEDDING_CONFIG_CACHE_TTL: Final = 60
|
||||
_EMBEDDING_CONFIG_CACHE_MAX_SIZE: Final = 256
|
||||
_embedding_config_cache: InMemoryCache | None = None
|
||||
|
||||
|
||||
def _get_embedding_config_cache() -> InMemoryCache:
|
||||
global _embedding_config_cache
|
||||
if _embedding_config_cache is None:
|
||||
_embedding_config_cache = InMemoryCache(
|
||||
max_size_in_memory=_EMBEDDING_CONFIG_CACHE_MAX_SIZE,
|
||||
default_ttl=_EMBEDDING_CONFIG_CACHE_TTL,
|
||||
)
|
||||
return _embedding_config_cache
|
||||
|
||||
|
||||
def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> Any:
|
||||
"""
|
||||
|
|
@ -155,235 +127,6 @@ async def _fetch_and_authorize_vector_store(
|
|||
return typed
|
||||
|
||||
|
||||
def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from router's config-defined models.
|
||||
|
||||
Config-defined models (from proxy_config.yaml) are stored in the router's model_list,
|
||||
not in the database. This function looks up the model in the router and extracts
|
||||
api_key, api_base, and api_version from the deployment's litellm_params.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
llm_router: The LiteLLM router instance
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model or llm_router is None:
|
||||
return None
|
||||
|
||||
# Extract model name candidates - could be "text-embedding-ada-002" or "azure/text-embedding-3-large"
|
||||
# Try exact match first, then try without provider prefix
|
||||
model_name_candidates: Final = [embedding_model]
|
||||
if "/" in embedding_model:
|
||||
# If it has a provider prefix, also try without it
|
||||
_, model_name = embedding_model.split("/", 1)
|
||||
model_name_candidates.append(model_name)
|
||||
|
||||
# Try to find model in router
|
||||
for model_name in model_name_candidates:
|
||||
try:
|
||||
# Try to get deployment by model group name (model_name in config)
|
||||
deployment = llm_router.get_deployment_by_model_group_name(model_group_name=model_name)
|
||||
|
||||
if deployment is not None and deployment.litellm_params is not None:
|
||||
litellm_params = deployment.litellm_params
|
||||
|
||||
# Build embedding config from model params
|
||||
embedding_config: dict[str, object] = {}
|
||||
|
||||
# Extract api_key
|
||||
api_key = getattr(litellm_params, "api_key", None)
|
||||
if api_key:
|
||||
# Handle os.environ/ prefix
|
||||
if isinstance(api_key, str) and api_key.startswith("os.environ/"):
|
||||
api_key = get_secret(api_key)
|
||||
embedding_config["api_key"] = api_key
|
||||
|
||||
# Extract api_base
|
||||
api_base = getattr(litellm_params, "api_base", None)
|
||||
if api_base:
|
||||
# Handle os.environ/ prefix
|
||||
if isinstance(api_base, str) and api_base.startswith("os.environ/"):
|
||||
api_base = get_secret(api_base)
|
||||
embedding_config["api_base"] = api_base
|
||||
|
||||
# Extract api_version
|
||||
api_version = getattr(litellm_params, "api_version", None)
|
||||
if api_version:
|
||||
embedding_config["api_version"] = api_version
|
||||
|
||||
project_id = getattr(litellm_params, "project_id", None)
|
||||
if project_id:
|
||||
embedding_config["project_id"] = project_id
|
||||
|
||||
# Only return config if we have at least api_key or api_base
|
||||
if embedding_config:
|
||||
verbose_proxy_logger.debug(
|
||||
"Resolved embedding config from router model %s: %s", model_name, list(embedding_config.keys())
|
||||
)
|
||||
return embedding_config
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Error resolving embedding config from router for model %s: %s", model_name, e)
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_embedding_config_from_db(
|
||||
embedding_model: str, prisma_client: "PrismaClient"
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from database model configuration.
|
||||
|
||||
If litellm_embedding_model is provided but litellm_embedding_config is not,
|
||||
this function looks up the model in the database and extracts api_key, api_base,
|
||||
and api_version from the model's litellm_params to build the embedding config.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model:
|
||||
return None
|
||||
|
||||
# Extract model name - could be "text-embedding-ada-002" or "azure/text-embedding-3-large"
|
||||
# Try to find model by exact match first, then try without provider prefix
|
||||
model_name_candidates: Final = [embedding_model]
|
||||
if "/" in embedding_model:
|
||||
# If it has a provider prefix, also try without it
|
||||
_, model_name = embedding_model.split("/", 1)
|
||||
model_name_candidates.append(model_name)
|
||||
|
||||
# Try to find model in database
|
||||
for model_name in model_name_candidates:
|
||||
try:
|
||||
db_model = await ModelRepository(prisma_client).table.find_first(where={"model_name": model_name})
|
||||
|
||||
if db_model and db_model.litellm_params:
|
||||
# Extract litellm_params (could be dict or JSON string)
|
||||
model_params = db_model.litellm_params
|
||||
if isinstance(model_params, str): # pyright: ignore[reportUnnecessaryIsInstance] # prisma Json is str
|
||||
model_params = json.loads(model_params)
|
||||
|
||||
# Decrypt values from database (similar to how proxy_server.py does it)
|
||||
# Values stored in DB are encrypted, so we need to decrypt them first
|
||||
decrypted_params = {}
|
||||
if isinstance(model_params, dict):
|
||||
for k, v in model_params.items():
|
||||
if isinstance(v, str):
|
||||
# Decrypt value - returns original value if decryption fails or no key is set
|
||||
decrypted_value = decrypt_value_helper(value=v, key=k, return_original_value=True)
|
||||
decrypted_params[k] = decrypted_value
|
||||
else:
|
||||
decrypted_params[k] = v
|
||||
else:
|
||||
decrypted_params = model_params
|
||||
|
||||
# Build embedding config from model params
|
||||
embedding_config = {}
|
||||
|
||||
# Extract api_key
|
||||
api_key = decrypted_params.get("api_key")
|
||||
if api_key:
|
||||
# Handle os.environ/ prefix (after decryption, values may be os.environ/ prefixed)
|
||||
if isinstance(api_key, str) and api_key.startswith("os.environ/"):
|
||||
api_key = get_secret(api_key)
|
||||
embedding_config["api_key"] = api_key
|
||||
|
||||
# Extract api_base
|
||||
api_base = decrypted_params.get("api_base")
|
||||
if api_base:
|
||||
# Handle os.environ/ prefix (after decryption, values may be os.environ/ prefixed)
|
||||
if isinstance(api_base, str) and api_base.startswith("os.environ/"):
|
||||
api_base = get_secret(api_base)
|
||||
embedding_config["api_base"] = api_base
|
||||
|
||||
# Extract api_version
|
||||
api_version = decrypted_params.get("api_version")
|
||||
if api_version:
|
||||
embedding_config["api_version"] = api_version
|
||||
|
||||
# Only return config if we have at least api_key or api_base
|
||||
if embedding_config:
|
||||
verbose_proxy_logger.debug(
|
||||
"Resolved embedding config from database model %s: %s",
|
||||
model_name,
|
||||
list(embedding_config.keys()),
|
||||
)
|
||||
return embedding_config
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Error resolving embedding config for model %s: %s", model_name, e)
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_embedding_config(
|
||||
embedding_model: str, prisma_client: "PrismaClient | None", llm_router: "Router | None" = None
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from either router (config-defined) or database models.
|
||||
|
||||
This function first checks the router for config-defined models, then falls back
|
||||
to the database. This allows users to use models defined in either location.
|
||||
|
||||
Results are cached in process memory for ``_EMBEDDING_CONFIG_CACHE_TTL``
|
||||
seconds so the request-handling path doesn't hit the database on every
|
||||
vector-store call. Negative results (model not found) are intentionally
|
||||
not cached to avoid blocking a freshly-added model behind the TTL.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
prisma_client: The Prisma client instance
|
||||
llm_router: The LiteLLM router instance (optional, will be imported if not provided)
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model:
|
||||
return None
|
||||
|
||||
cache: Final = _get_embedding_config_cache()
|
||||
cached: Final = cache.get_cache(embedding_model)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# Import llm_router if not provided
|
||||
if llm_router is None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
llm_router = None
|
||||
|
||||
# First try to resolve from router (config-defined models)
|
||||
if llm_router is not None:
|
||||
router_config = _resolve_embedding_config_from_router(embedding_model=embedding_model, llm_router=llm_router)
|
||||
if router_config:
|
||||
verbose_proxy_logger.debug("Resolved embedding config from router for model %s", embedding_model)
|
||||
cache.set_cache(embedding_model, router_config)
|
||||
return router_config
|
||||
|
||||
# Fall back to database
|
||||
if prisma_client is not None:
|
||||
db_config: Final = await _resolve_embedding_config_from_db(
|
||||
embedding_model=embedding_model, prisma_client=prisma_client
|
||||
)
|
||||
if db_config:
|
||||
verbose_proxy_logger.debug("Resolved embedding config from database for model %s", embedding_model)
|
||||
cache.set_cache(embedding_model, db_config)
|
||||
return db_config
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Could not resolve embedding config for model %s from router or database", embedding_model
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
########################################################
|
||||
# Helper Functions
|
||||
########################################################
|
||||
|
|
@ -469,10 +212,9 @@ async def create_vector_store_in_db(
|
|||
# (``api_key``, ``api_base``, ``api_version``) into this row. That
|
||||
# exposed every env-stored embedding-model credential on the
|
||||
# ``/vector_store/{new,info,update,list}`` responses. Keep the user's
|
||||
# raw ``litellm_embedding_model`` reference; resolution now happens in
|
||||
# ``build_request_data_from_managed_vector_store``
|
||||
# at request-handling time so the cleartext config exists only in
|
||||
# per-request memory and never reaches the database.
|
||||
# raw ``litellm_embedding_model`` reference; each search embeds the
|
||||
# query through the router at request time, so the credentials stay
|
||||
# on the deployment and never reach the database.
|
||||
if litellm_params:
|
||||
litellm_params_dict: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True)
|
||||
data_to_create["litellm_params"] = safe_dumps(litellm_params_dict)
|
||||
|
|
@ -862,11 +604,9 @@ async def update_vector_store(
|
|||
|
||||
# Handle litellm_params if provided. As with the create path, the
|
||||
# embedding-config auto-resolve previously persisted cleartext
|
||||
# credentials into the row; resolution now happens at request-
|
||||
# handling time in
|
||||
# ``build_request_data_from_managed_vector_store``
|
||||
# so this row only ever stores the user-supplied
|
||||
# ``litellm_embedding_model`` reference.
|
||||
# credentials into the row; each search now embeds the query
|
||||
# through the router at request time, so this row only ever stores
|
||||
# the user-supplied ``litellm_embedding_model`` reference.
|
||||
if "litellm_params" in update_data:
|
||||
_input_litellm_params: Final[dict] = update_data.get("litellm_params", {}) or {}
|
||||
litellm_params_dict: Final = GenericLiteLLMParams(**_input_litellm_params).model_dump(exclude_none=True)
|
||||
|
|
|
|||
|
|
@ -85,6 +85,10 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
|
|||
mask_credentials_in_payload,
|
||||
mask_sensitive_structure,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
vector_store_request_metadata,
|
||||
)
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
|
||||
|
|
@ -6481,11 +6485,24 @@ class Router:
|
|||
if custom_llm_provider and "custom_llm_provider" not in kwargs
|
||||
else MappingProxyType(kwargs)
|
||||
)
|
||||
if provider_kwargs.get("model"):
|
||||
return self._generic_api_call_with_fallbacks(original_function=original_function, **provider_kwargs)
|
||||
search_kwargs: Final = (
|
||||
MappingProxyType(
|
||||
{
|
||||
**provider_kwargs,
|
||||
"_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor(
|
||||
router=self,
|
||||
metadata=self._vector_store_request_metadata(kwargs),
|
||||
),
|
||||
}
|
||||
)
|
||||
if call_type == "vector_store_search"
|
||||
else provider_kwargs
|
||||
)
|
||||
if search_kwargs.get("model"):
|
||||
return self._generic_api_call_with_fallbacks(original_function=original_function, **search_kwargs)
|
||||
if call_type == "vector_store_search":
|
||||
return original_function(**MappingProxyType({**provider_kwargs, "router": self}))
|
||||
return original_function(**provider_kwargs)
|
||||
return original_function(**MappingProxyType({**search_kwargs, "router": self}))
|
||||
return original_function(**search_kwargs)
|
||||
|
||||
return vector_store_sync_wrapper
|
||||
|
||||
|
|
@ -6658,11 +6675,22 @@ class Router:
|
|||
"avector_store_update",
|
||||
"avector_store_delete",
|
||||
):
|
||||
vector_store_kwargs: Final = (
|
||||
{ # mutable-ok: the async routed request requires dynamic keyword arguments
|
||||
**kwargs,
|
||||
"_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor(
|
||||
router=self,
|
||||
metadata=self._vector_store_request_metadata(kwargs),
|
||||
),
|
||||
}
|
||||
if call_type == "avector_store_search"
|
||||
else kwargs
|
||||
)
|
||||
return await self._init_vector_store_api_endpoints(
|
||||
original_function=original_function,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
call_type=call_type,
|
||||
**kwargs,
|
||||
**vector_store_kwargs,
|
||||
)
|
||||
elif call_type in ("afile_delete", "afile_content"):
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
|
|
@ -6698,6 +6726,10 @@ class Router:
|
|||
|
||||
return async_wrapper
|
||||
|
||||
@staticmethod
|
||||
def _vector_store_request_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return vector_store_request_metadata(kwargs)
|
||||
|
||||
async def _init_vector_store_api_endpoints(
|
||||
self,
|
||||
original_function: Callable,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,11 @@ import litellm
|
|||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
vector_store_request_metadata,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -38,6 +43,16 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
#################################################
|
||||
|
||||
|
||||
def _direct_vector_store_embedding_executor(
|
||||
value: object, router: "Router | None", request_kwargs: Mapping[str, object]
|
||||
) -> VectorStoreEmbeddingExecutor:
|
||||
if value is not None and not isinstance(value, VectorStoreEmbeddingExecutor):
|
||||
raise TypeError("Invalid direct vector store embedding executor")
|
||||
return BaseQueryEmbeddingVectorStoreConfig.query_embedding_executor(
|
||||
value, router, vector_store_request_metadata(request_kwargs)
|
||||
)
|
||||
|
||||
|
||||
def mock_vector_store_search_response(
|
||||
mock_results: list[VectorStoreSearchResult] | None = None,
|
||||
):
|
||||
|
|
@ -289,7 +304,12 @@ async def asearch(
|
|||
"""
|
||||
Async: Search a vector store for relevant chunks based on a query and file attributes filter.
|
||||
"""
|
||||
local_vars: Final = locals()
|
||||
embedding_executor: Final = _direct_vector_store_embedding_executor(
|
||||
kwargs.pop("_direct_vector_store_embedding_executor", None), router, kwargs
|
||||
)
|
||||
local_vars: Final = { # mutable-ok: exception logging requires a sanitized mutable snapshot
|
||||
key: value for key, value in locals().items() if key != "embedding_executor"
|
||||
}
|
||||
|
||||
try:
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
|
|
@ -312,6 +332,7 @@ async def asearch(
|
|||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
_direct_vector_store_embedding_executor=embedding_executor,
|
||||
router=router,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -369,12 +390,16 @@ def search(
|
|||
Returns:
|
||||
VectorStoreSearchResponse containing the search results.
|
||||
"""
|
||||
local_vars: Final = locals()
|
||||
embedding_executor: Final = _direct_vector_store_embedding_executor(
|
||||
kwargs.pop("_direct_vector_store_embedding_executor", None), router, kwargs
|
||||
)
|
||||
local_vars: Final = { # mutable-ok: exception logging requires a sanitized mutable snapshot
|
||||
key: value for key, value in locals().items() if key != "embedding_executor"
|
||||
}
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("asearch", False) is True
|
||||
|
||||
# pull credentials from registry if available
|
||||
if litellm.vector_store_registry is not None and vector_store_id is not None:
|
||||
try:
|
||||
|
|
@ -451,6 +476,7 @@ def search(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
embedding_executor=embedding_executor,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout or request_timeout,
|
||||
|
|
|
|||
|
|
@ -71,6 +71,48 @@ def setup_vector_store_registry():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_hook_routes_search_through_proxy_router(
|
||||
setup_vector_store_registry,
|
||||
):
|
||||
proxy_router = Mock()
|
||||
proxy_router.avector_store_search = AsyncMock(
|
||||
return_value=VectorStoreSearchResponse(
|
||||
object="vector_store.search_results.page",
|
||||
search_query="what is litellm?",
|
||||
data=[
|
||||
VectorStoreSearchResult(
|
||||
score=1.0,
|
||||
content=[VectorStoreResultContent(text="routed context", type="text")],
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
logging_obj = Mock()
|
||||
logging_obj.model_call_details = {
|
||||
"litellm_params": {"metadata": {"user_api_key_team_id": "team-a"}}
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", proxy_router):
|
||||
_, messages, _ = await VectorStorePreCallHook().async_get_chat_completion_prompt(
|
||||
model="chat-model",
|
||||
messages=[{"role": "user", "content": "what is litellm?"}],
|
||||
non_default_params={"vector_store_ids": ["T37J8R4WTM"]},
|
||||
prompt_id=None,
|
||||
prompt_variables=None,
|
||||
dynamic_callback_params={},
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
proxy_router.avector_store_search.assert_awaited_once_with(
|
||||
vector_store_id="T37J8R4WTM",
|
||||
query="what is litellm?",
|
||||
custom_llm_provider="bedrock",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
assert messages[0]["content"] == "Context:\n\nrouted context\n\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(
|
||||
setup_vector_store_registry,
|
||||
|
|
|
|||
|
|
@ -5,17 +5,218 @@ These tests simulate real-world scenarios where headers and configuration
|
|||
need to be properly propagated through the router to the LLM API.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch, AsyncMock
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
|
||||
QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
OPENAI_EMBEDDINGS_URL = "https://api.openai.com/v1/embeddings"
|
||||
STORE_EMBEDDINGS_URL = "https://embedding.example/v1/embeddings"
|
||||
|
||||
|
||||
def _mock_embedding_route(respx_mock: respx.MockRouter, url: str) -> respx.Route:
|
||||
return respx_mock.post(url).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _sent(route: respx.Route, index: int) -> tuple[str, str, list[str]]:
|
||||
request = route.calls[index].request
|
||||
body = json.loads(request.read())
|
||||
return request.headers["authorization"], body["model"], body["input"]
|
||||
|
||||
|
||||
def _alias_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "team-alias",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class TestRouterEmbeddingIntegration:
|
||||
"""Integration tests for embedding with router configuration."""
|
||||
|
||||
def test_vector_store_request_metadata_prefers_litellm_metadata(self):
|
||||
assert Router._vector_store_request_metadata(
|
||||
{
|
||||
"litellm_metadata": {"user_api_key_team_id": "team-a"},
|
||||
"metadata": {"user_api_key_team_id": "team-b"},
|
||||
}
|
||||
) == {"user_api_key_team_id": "team-a"}
|
||||
|
||||
assert Router._vector_store_request_metadata({"metadata": {"user_api_key_team_id": "team-b"}}) == {
|
||||
"user_api_key_team_id": "team-b"
|
||||
}
|
||||
assert Router._vector_store_request_metadata({}) == {}
|
||||
|
||||
def test_sync_vector_store_wrapper_injects_router_embedding_executor(self):
|
||||
router = Router(model_list=[])
|
||||
original = MagicMock(return_value="searched")
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
assert (
|
||||
wrapped(
|
||||
vector_store_id="store",
|
||||
query="query",
|
||||
custom_llm_provider="valkey",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
== "searched"
|
||||
)
|
||||
|
||||
call_kwargs = original.call_args.kwargs
|
||||
assert call_kwargs["custom_llm_provider"] == "valkey"
|
||||
executor = call_kwargs["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata == {"user_api_key_team_id": "team-a"}
|
||||
|
||||
def test_sync_vector_store_wrapper_preserves_model_routing(self):
|
||||
router = Router(model_list=[])
|
||||
original = MagicMock()
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
with patch.object(router, "_generic_api_call_with_fallbacks", return_value="routed") as fallback:
|
||||
assert wrapped(model="vector-alias", vector_store_id="store", query="query") == "routed"
|
||||
|
||||
assert fallback.call_args.kwargs["model"] == "vector-alias"
|
||||
assert fallback.call_args.kwargs["original_function"] is original
|
||||
assert isinstance(
|
||||
fallback.call_args.kwargs["_direct_vector_store_embedding_executor"],
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executors_cover_sdk_and_router_paths(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
store_route = _mock_embedding_route(respx_mock, STORE_EMBEDDINGS_URL)
|
||||
sdk_executor = LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
sync_response = sdk_executor.embed("openai/text-embedding-3-small", "sync", {"api_key": "explicit"})
|
||||
async_response = await sdk_executor.aembed("openai/text-embedding-3-small", "async", {"api_key": "explicit"})
|
||||
|
||||
assert sync_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(openai_route, 0) == ("Bearer explicit", "text-embedding-3-small", ["sync"])
|
||||
assert _sent(openai_route, 1) == ("Bearer explicit", "text-embedding-3-small", ["async"])
|
||||
|
||||
explicit_config = {
|
||||
"api_base": "https://embedding.example/v1",
|
||||
"api_key": "store-key",
|
||||
"metadata": {
|
||||
"configured": True,
|
||||
"user_api_key_team_id": "untrusted-team",
|
||||
},
|
||||
"model": "untrusted-model",
|
||||
}
|
||||
mock_router = MagicMock()
|
||||
mock_router.embedding.return_value = sync_response
|
||||
router_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=mock_router,
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
assert router_executor.embed("team-alias", "query", explicit_config) is sync_response
|
||||
mock_router.embedding.assert_called_once_with(
|
||||
model="team-alias",
|
||||
input=["query"],
|
||||
api_base="https://embedding.example/v1",
|
||||
api_key="store-key",
|
||||
metadata={"configured": True, "user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
alias_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
sync_alias = alias_executor.embed("team-alias", "sync query", explicit_config)
|
||||
async_alias = await alias_executor.aembed("team-alias", "async query", explicit_config)
|
||||
|
||||
assert sync_alias.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_alias.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert openai_route.call_count == 2
|
||||
assert _sent(store_route, 0) == ("Bearer store-key", "text-embedding-3-small", ["sync query"])
|
||||
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-small", ["async query"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_executor_falls_back_to_sdk_for_models_the_router_does_not_serve(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
store_route = _mock_embedding_route(respx_mock, STORE_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
inline_config = {"api_base": "https://embedding.example/v1", "api_key": "store-key"}
|
||||
|
||||
sync_response = executor.embed("openai/text-embedding-3-large", "sync query", inline_config)
|
||||
async_response = await executor.aembed("openai/text-embedding-3-large", "async query", inline_config)
|
||||
|
||||
assert sync_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(store_route, 0) == ("Bearer store-key", "text-embedding-3-large", ["sync query"])
|
||||
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-large", ["async query"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_executor_rejects_unserved_models_without_explicit_config(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "env-key")
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
executor.embed("openai/text-embedding-3-large", "sync query", {})
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
await executor.aembed("openai/text-embedding-3-large", "async query", {})
|
||||
|
||||
assert openai_route.call_count == 0
|
||||
|
||||
def test_router_executor_routes_deployment_model_names_through_the_router(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(router=_alias_router(), metadata={})
|
||||
|
||||
response = executor.embed("openai/text-embedding-3-small", "query", {})
|
||||
|
||||
assert response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(openai_route, 0) == ("Bearer deployment-key", "text-embedding-3-small", ["query"])
|
||||
|
||||
def test_embedding_with_deployment_specific_headers(self):
|
||||
"""
|
||||
Test that deployment-specific headers are propagated.
|
||||
|
|
@ -122,9 +323,7 @@ class TestRouterEmbeddingIntegration:
|
|||
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
default_litellm_params={
|
||||
"metadata": {"environment": "test", "service": "embedding-service"}
|
||||
},
|
||||
default_litellm_params={"metadata": {"environment": "test", "service": "embedding-service"}},
|
||||
)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
|
|
@ -240,9 +439,7 @@ class TestRouterEmbeddingIntegration:
|
|||
# Make multiple calls and verify headers are always present
|
||||
for i in range(5):
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MagicMock(
|
||||
data=[{"embedding": [0.1, 0.2]}]
|
||||
)
|
||||
mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}])
|
||||
|
||||
router.embedding(model="shared-embedding-model", input=[f"test {i}"])
|
||||
|
||||
|
|
@ -327,9 +524,7 @@ class TestRouterEmbeddingIntegration:
|
|||
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
default_litellm_params={
|
||||
"headers": {"X-Custom-Azure-Header": "azure-value"}
|
||||
},
|
||||
default_litellm_params={"headers": {"X-Custom-Azure-Header": "azure-value"}},
|
||||
)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
|
|
|
|||
|
|
@ -67,20 +67,52 @@ class FakeAsyncEmbeddingFn(FakeEmbeddingFn):
|
|||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
|
||||
class FakeEmbeddingExecutor:
|
||||
def __init__(self, embedding):
|
||||
self.embedding = embedding
|
||||
self.captured = None
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.captured = (model, query, configuration)
|
||||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.captured = (model, query, configuration)
|
||||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
|
||||
def _doc(doc_id, distance, **fields):
|
||||
return SimpleNamespace(id=doc_id, vector_distance=str(distance), **fields)
|
||||
|
||||
|
||||
def _search(config, client=None, query="what is litellm", optional_params=None, litellm_params=None):
|
||||
def _search(config, client=None, query="what is litellm", optional_params=None, litellm_params=None, executor=None):
|
||||
return config.execute_search_vector_store_request(
|
||||
vector_store_id="my_index",
|
||||
query=query,
|
||||
vector_store_search_optional_params=optional_params or {},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small", **(litellm_params or {})},
|
||||
embedding_executor=executor,
|
||||
)
|
||||
|
||||
|
||||
def test_sync_search_uses_request_embedding_executor_without_overwriting_explicit_config():
|
||||
executor = FakeEmbeddingExecutor([0.1, 0.2])
|
||||
config = ValkeyVectorStoreConfig(sync_client=FakeRedis())
|
||||
embedding_config = {"api_key": "store-specific-key", "aws_region_name": "us-west-2"}
|
||||
|
||||
_search(
|
||||
config,
|
||||
litellm_params={
|
||||
"litellm_embedding_model": "team-embedding-alias",
|
||||
"litellm_embedding_config": embedding_config,
|
||||
},
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
assert executor.captured == ("team-embedding-alias", "what is litellm", embedding_config)
|
||||
|
||||
|
||||
def test_sync_search_builds_knn_query_with_packed_vector():
|
||||
embedding_fn = FakeEmbeddingFn([0.1, 0.2, 0.3])
|
||||
client = FakeRedis()
|
||||
|
|
|
|||
|
|
@ -2,29 +2,24 @@ from datetime import datetime, timezone
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import (
|
||||
_update_request_data_with_litellm_managed_vector_store_registry,
|
||||
index_create,
|
||||
index_list,
|
||||
)
|
||||
from litellm.proxy.vector_store_files_endpoints.endpoints import (
|
||||
_update_request_data_with_model_routing_hint,
|
||||
)
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_check_vector_store_access,
|
||||
_resolve_embedding_config,
|
||||
_resolve_embedding_config_from_db,
|
||||
_resolve_embedding_config_from_router,
|
||||
create_vector_store_in_db,
|
||||
new_vector_store,
|
||||
)
|
||||
|
|
@ -33,8 +28,12 @@ from litellm.proxy.vector_store_endpoints.utils import (
|
|||
is_allowed_to_call_vector_store_endpoint,
|
||||
is_allowed_to_call_vector_store_files_endpoint,
|
||||
)
|
||||
from litellm.proxy.vector_store_files_endpoints.endpoints import (
|
||||
_update_request_data_with_model_routing_hint,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse, LlmProviders
|
||||
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.vector_stores.main import _direct_vector_store_embedding_executor
|
||||
|
||||
|
||||
def _serialize_litellm_params(litellm_params):
|
||||
|
|
@ -51,17 +50,113 @@ def _serialize_litellm_params(litellm_params):
|
|||
return json.dumps(litellm_params or {})
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_embedding_config_cache():
|
||||
"""The use-time embedding-config resolver caches results in process
|
||||
memory across calls. Reset it before every test so the resolver
|
||||
actually exercises the router/DB path under test instead of returning
|
||||
a value cached by an earlier test."""
|
||||
from litellm.proxy.vector_store_endpoints import management_endpoints
|
||||
def test_direct_vector_store_embedding_executor_rejects_invalid_value():
|
||||
with pytest.raises(TypeError, match="Invalid direct vector store embedding executor"):
|
||||
_direct_vector_store_embedding_executor(object(), None, {})
|
||||
|
||||
management_endpoints._embedding_config_cache = None
|
||||
yield
|
||||
management_endpoints._embedding_config_cache = None
|
||||
|
||||
def test_router_vector_store_search_injects_executor_and_request_metadata():
|
||||
router = litellm.Router(model_list=[])
|
||||
original = MagicMock(return_value="searched")
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
assert (
|
||||
wrapped(
|
||||
vector_store_id="store",
|
||||
query="query",
|
||||
custom_llm_provider="valkey",
|
||||
litellm_metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
== "searched"
|
||||
)
|
||||
|
||||
call_kwargs = original.call_args.kwargs
|
||||
assert call_kwargs["custom_llm_provider"] == "valkey"
|
||||
executor = call_kwargs["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata == {"user_api_key_team_id": "team-a"}
|
||||
assert litellm.Router._vector_store_request_metadata({"metadata": {"user_api_key_team_id": "team-b"}}) == {
|
||||
"user_api_key_team_id": "team-b"
|
||||
}
|
||||
assert litellm.Router._vector_store_request_metadata({}) == {}
|
||||
|
||||
with patch.object( # test-quality-ok: fallback dispatch is the boundary this wrapper delegates to
|
||||
router, "_generic_api_call_with_fallbacks", return_value="routed"
|
||||
) as fallback:
|
||||
assert wrapped(model="vector-alias", vector_store_id="store", query="query") == "routed"
|
||||
assert fallback.call_args.kwargs["model"] == "vector-alias"
|
||||
assert fallback.call_args.kwargs["original_function"] is original
|
||||
|
||||
create_original = MagicMock(return_value="created")
|
||||
wrapped_create = router.factory_function(create_original, call_type="vector_store_create")
|
||||
assert wrapped_create(name="store") == "created"
|
||||
create_original.assert_called_once_with(name="store")
|
||||
with patch.object( # test-quality-ok: fallback dispatch is the boundary this wrapper delegates to
|
||||
router, "_generic_api_call_with_fallbacks", return_value="created-through-router"
|
||||
) as fallback:
|
||||
assert wrapped_create(model="vector-alias", name="store") == "created-through-router"
|
||||
fallback.assert_called_once_with(original_function=create_original, model="vector-alias", name="store")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executors_preserve_explicit_configuration():
|
||||
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
|
||||
sdk_executor = LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: isolates SDK dispatch from external embedding providers
|
||||
"litellm.embedding", return_value=response
|
||||
) as embedding,
|
||||
patch( # test-quality-ok: isolates async SDK dispatch from external embedding providers
|
||||
"litellm.aembedding", new=AsyncMock(return_value=response)
|
||||
) as aembedding,
|
||||
):
|
||||
assert sdk_executor.embed("openai/model", "sync", {"api_key": "explicit"}) is response
|
||||
assert await sdk_executor.aembed("openai/model", "async", {"api_key": "explicit"}) is response
|
||||
|
||||
embedding.assert_called_once_with(model="openai/model", input=["sync"], api_key="explicit")
|
||||
aembedding.assert_awaited_once_with(model="openai/model", input=["async"], api_key="explicit")
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.embedding.return_value = response
|
||||
mock_router.aembedding = AsyncMock(return_value=response)
|
||||
router_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=mock_router,
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
assert router_executor.embed("team-alias", "query", {}) is response
|
||||
mock_router.embedding.assert_called_once_with(
|
||||
model="team-alias",
|
||||
input=["query"],
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: verifies explicit store configuration at the SDK boundary
|
||||
"litellm.embedding", return_value=response
|
||||
) as explicit_embedding,
|
||||
patch( # test-quality-ok: verifies async explicit store configuration at the SDK boundary
|
||||
"litellm.aembedding", new=AsyncMock(return_value=response)
|
||||
) as explicit_aembedding,
|
||||
):
|
||||
assert router_executor.embed("openai/model", "query", {"api_key": "store-key"}) is response
|
||||
assert await router_executor.aembed("openai/model", "query", {"api_key": "store-key"}) is response
|
||||
|
||||
explicit_embedding.assert_not_called()
|
||||
explicit_aembedding.assert_not_awaited()
|
||||
assert mock_router.embedding.call_args.kwargs == {
|
||||
"model": "openai/model",
|
||||
"input": ["query"],
|
||||
"api_key": "store-key",
|
||||
"metadata": {"user_api_key_team_id": "team-a"},
|
||||
}
|
||||
mock_router.aembedding.assert_awaited_once_with(
|
||||
model="openai/model",
|
||||
input=["query"],
|
||||
api_key="store-key",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -82,10 +177,11 @@ async def test_router_avector_store_search_passes_correct_args():
|
|||
}
|
||||
|
||||
# Call router's avector_store_search
|
||||
result = await router.avector_store_search(
|
||||
await router.avector_store_search(
|
||||
vector_store_id="test_store_id",
|
||||
query="test query",
|
||||
custom_llm_provider="bedrock",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
# Verify the internal method was called with correct args
|
||||
|
|
@ -96,6 +192,38 @@ async def test_router_avector_store_search_passes_correct_args():
|
|||
assert call_args[1]["vector_store_id"] == "test_store_id"
|
||||
assert call_args[1]["query"] == "test query"
|
||||
assert call_args[1]["custom_llm_provider"] == "bedrock"
|
||||
executor = call_args[1]["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata["user_api_key_team_id"] == "team-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executor_uses_team_scoped_router_deployment():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "shared-embedding",
|
||||
"litellm_params": {"model": "openai/text-embedding-3-small", "api_key": "team-a-key"},
|
||||
"model_info": {"team_id": "team-a", "team_public_model_name": "shared-embedding"},
|
||||
},
|
||||
{
|
||||
"model_name": "shared-embedding",
|
||||
"litellm_params": {"model": "openai/text-embedding-3-small", "api_key": "team-b-key"},
|
||||
"model_info": {"team_id": "team-b", "team_public_model_name": "shared-embedding"},
|
||||
},
|
||||
]
|
||||
)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=router,
|
||||
metadata={"user_api_key_team_id": "team-b"},
|
||||
)
|
||||
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
|
||||
|
||||
with patch("litellm.aembedding", new=AsyncMock(return_value=response)) as mock_aembedding:
|
||||
result = await executor.aembed("shared-embedding", "query", {})
|
||||
|
||||
assert result is response
|
||||
assert mock_aembedding.await_args.kwargs["api_key"] == "team-b-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -502,91 +630,30 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_request_data_resolves_embedding_config_at_use_time():
|
||||
"""When the persisted vector store row carries only a
|
||||
``litellm_embedding_model`` reference (the new behaviour after
|
||||
moving the auto-resolve out of write time), the request-handling
|
||||
layer must resolve the embedding config so the downstream embed
|
||||
call still has ``api_key`` / ``api_base`` / ``api_version``. The
|
||||
resolved config lives in this per-request data dict only — never
|
||||
persisted."""
|
||||
mock_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
async def test_managed_vector_store_keeps_embedding_reference_and_explicit_config():
|
||||
explicit_config = {"api_key": "store-specific-key", "api_base": "https://embedding.example"}
|
||||
managed_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "test_store",
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"custom_llm_provider": "valkey",
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "azure/text-embedding-3-large",
|
||||
# Note: no litellm_embedding_config persisted
|
||||
"litellm_embedding_model": "team-embedding-alias",
|
||||
"litellm_embedding_config": explicit_config,
|
||||
},
|
||||
}
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = (
|
||||
mock_vector_store
|
||||
)
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
|
||||
|
||||
resolved = {
|
||||
"api_key": "use-time-resolved-key",
|
||||
"api_base": "https://my-azure.example",
|
||||
"api_version": "2024-09-01",
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints._resolve_embedding_config",
|
||||
new=AsyncMock(return_value=resolved),
|
||||
),
|
||||
):
|
||||
with patch.object(litellm, "vector_store_registry", mock_registry):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={}, vector_store_id="test_store"
|
||||
data={},
|
||||
vector_store_id="test_store",
|
||||
)
|
||||
|
||||
assert result["litellm_embedding_model"] == "azure/text-embedding-3-large"
|
||||
assert result["litellm_embedding_config"] == resolved
|
||||
assert result["litellm_embedding_model"] == "team-embedding-alias"
|
||||
assert result["litellm_embedding_config"] == explicit_config
|
||||
assert managed_vector_store["litellm_params"]["litellm_embedding_config"] == explicit_config
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_request_data_passes_through_legacy_embedding_config():
|
||||
"""A vector store row created by an older proxy version may already
|
||||
carry a fully-resolved ``litellm_embedding_config`` in its persisted
|
||||
``litellm_params`` (the very leak this PR closes). Those legacy rows
|
||||
must still work — the use-time resolver skips re-resolution when
|
||||
the config is already present so the embed call keeps succeeding."""
|
||||
legacy_config = {
|
||||
"api_key": "legacy-cleartext-key",
|
||||
"api_base": "https://legacy-azure.example",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
mock_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "legacy_store",
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "azure/text-embedding-3-large",
|
||||
"litellm_embedding_config": legacy_config,
|
||||
},
|
||||
}
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = (
|
||||
mock_vector_store
|
||||
)
|
||||
|
||||
resolve_mock = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints._resolve_embedding_config",
|
||||
new=resolve_mock,
|
||||
),
|
||||
):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={}, vector_store_id="legacy_store"
|
||||
)
|
||||
|
||||
assert result["litellm_embedding_config"] == legacy_config
|
||||
resolve_mock.assert_not_awaited()
|
||||
|
||||
|
||||
class TestCheckVectorStorePermission:
|
||||
"""Test suite for check_vector_store_permission function."""
|
||||
|
|
@ -2003,57 +2070,7 @@ async def test_vector_store_update_and_list_synchronization():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_from_db():
|
||||
"""Test that _resolve_embedding_config_from_db correctly resolves embedding config from database."""
|
||||
mock_prisma_client = MagicMock()
|
||||
|
||||
# Mock database model with litellm_params
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://api.openai.com",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
):
|
||||
result = await _resolve_embedding_config_from_db(
|
||||
embedding_model="text-embedding-ada-002", prisma_client=mock_prisma_client
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "test-api-key"
|
||||
assert result["api_base"] == "https://api.openai.com"
|
||||
assert result["api_version"] == "2024-01-01"
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_called_once_with(
|
||||
where={"model_name": "text-embedding-ada-002"}
|
||||
)
|
||||
|
||||
# Test with empty embedding_model
|
||||
result_empty = await _resolve_embedding_config_from_db(
|
||||
embedding_model="", prisma_client=mock_prisma_client
|
||||
)
|
||||
assert result_empty is None
|
||||
|
||||
# Test with model not found
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
result_not_found = await _resolve_embedding_config_from_db(
|
||||
embedding_model="non-existent-model", prisma_client=mock_prisma_client
|
||||
)
|
||||
assert result_not_found is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_vector_store_auto_resolves_embedding_config():
|
||||
"""Test that new_vector_store auto-resolves embedding config when embedding_model is provided but config is not."""
|
||||
async def test_new_vector_store_persists_embedding_reference_without_credentials():
|
||||
import json
|
||||
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
|
|
@ -2070,14 +2087,6 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
},
|
||||
}
|
||||
|
||||
# Mock database model lookup for embedding config resolution
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "resolved-api-key",
|
||||
"api_base": "https://api.openai.com",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
|
||||
# Mock user API key
|
||||
mock_user_api_key = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_api_key.user_role = None
|
||||
|
|
@ -2088,10 +2097,6 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(
|
||||
return_value=None # Vector store doesn't exist yet
|
||||
)
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
# Track what was passed to create
|
||||
captured_create_data = {}
|
||||
|
||||
|
|
@ -2112,261 +2117,21 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
mock_registry = MagicMock()
|
||||
mock_registry.add_vector_store_to_registry = MagicMock()
|
||||
|
||||
# Mock router to return None (so it falls back to DB resolution)
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.proxy.proxy_server.llm_router", mock_router),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
),
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
):
|
||||
result = await new_vector_store(
|
||||
vector_store=vector_store_data, user_api_key_dict=mock_user_api_key
|
||||
)
|
||||
result = await new_vector_store(vector_store=vector_store_data, user_api_key_dict=mock_user_api_key)
|
||||
|
||||
assert result["status"] == "success"
|
||||
# Auto-resolve no longer happens at create time — the persisted row
|
||||
# carries only the model reference, never the resolved cleartext
|
||||
# credential. Resolution now happens at request-handling time inside
|
||||
# ``_update_request_data_with_litellm_managed_vector_store_registry``,
|
||||
# where the resolved config lives in per-request memory and is never
|
||||
# written to the database.
|
||||
litellm_params_json = captured_create_data.get("litellm_params")
|
||||
assert litellm_params_json is not None
|
||||
litellm_params_dict = json.loads(litellm_params_json)
|
||||
assert "litellm_embedding_config" not in litellm_params_dict
|
||||
assert litellm_params_dict["litellm_embedding_model"] == "text-embedding-ada-002"
|
||||
|
||||
# The response must also not echo a cleartext credential — even on
|
||||
# the create response, where redaction guards against caller-supplied
|
||||
# cleartext or pre-existing rows that were created by an earlier
|
||||
# proxy version.
|
||||
response_vs = result["vector_store"]
|
||||
assert "resolved-api-key" not in _serialize_litellm_params(
|
||||
response_vs.get("litellm_params")
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router():
|
||||
"""Test that _resolve_embedding_config_from_router correctly extracts credentials from config-defined models."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
# Create a mock router with a model
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Create a mock deployment with litellm_params
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "config-api-key"
|
||||
mock_litellm_params.api_base = "https://config-api-base.com"
|
||||
mock_litellm_params.api_version = "2024-02-01"
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
# Test resolution
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="text-embedding-ada-002", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "config-api-key"
|
||||
assert result["api_base"] == "https://config-api-base.com"
|
||||
assert result["api_version"] == "2024-02-01"
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.assert_called_once_with(
|
||||
model_group_name="text-embedding-ada-002"
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_with_provider_prefix():
|
||||
"""Test that _resolve_embedding_config_from_router handles provider prefixes like 'azure/model-name'."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
# Create a mock router
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Create a mock deployment
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "azure-api-key"
|
||||
mock_litellm_params.api_base = "https://azure-endpoint.openai.azure.com"
|
||||
mock_litellm_params.api_version = "2024-02-15"
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
# First call with full name returns None, second call with stripped name returns deployment
|
||||
mock_router.get_deployment_by_model_group_name.side_effect = [None, mock_deployment]
|
||||
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="azure/text-embedding-3-large", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "azure-api-key"
|
||||
assert result["api_base"] == "https://azure-endpoint.openai.azure.com"
|
||||
assert result["api_version"] == "2024-02-15"
|
||||
|
||||
# Should have tried both the full name and stripped name
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 2
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_returns_none_when_not_found():
|
||||
"""Test that _resolve_embedding_config_from_router returns None when model is not in router."""
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="nonexistent-model", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_handles_os_environ():
|
||||
"""Test that _resolve_embedding_config_from_router handles os.environ/ prefixed values."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_router = MagicMock()
|
||||
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "os.environ/OPENAI_API_KEY"
|
||||
mock_litellm_params.api_base = "https://direct-url.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.get_secret",
|
||||
return_value="resolved-from-env",
|
||||
) as mock_get_secret:
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="text-embedding-ada-002", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "resolved-from-env"
|
||||
assert result["api_base"] == "https://direct-url.com"
|
||||
assert "api_version" not in result
|
||||
|
||||
mock_get_secret.assert_called_once_with("os.environ/OPENAI_API_KEY")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_tries_router_then_db():
|
||||
"""Test that _resolve_embedding_config tries router first, then falls back to DB."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Router has the model
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "router-api-key"
|
||||
mock_litellm_params.api_base = "https://router-api-base.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
# DB should NOT be called since router has the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock()
|
||||
|
||||
result = await _resolve_embedding_config(
|
||||
embedding_model="text-embedding-ada-002",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "router-api-key"
|
||||
|
||||
# DB should NOT have been called since router found the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_caches_result():
|
||||
"""The first lookup should hit the router/DB; subsequent lookups for
|
||||
the same model name should return the cached value without touching
|
||||
the router or the database."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "router-api-key"
|
||||
mock_litellm_params.api_base = "https://router-api-base.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
first = await _resolve_embedding_config(
|
||||
embedding_model="cached-model",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
assert first is not None
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 1
|
||||
|
||||
second = await _resolve_embedding_config(
|
||||
embedding_model="cached-model",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
assert second == first
|
||||
# Router (and by extension the DB) was not consulted again.
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_falls_back_to_db():
|
||||
"""Test that _resolve_embedding_config falls back to DB when router doesn't have the model."""
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Router doesn't have the model
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
# DB has the model
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "db-api-key",
|
||||
"api_base": "https://db-api-base.com",
|
||||
}
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
):
|
||||
result = await _resolve_embedding_config(
|
||||
embedding_model="text-embedding-ada-002",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "db-api-key"
|
||||
|
||||
# DB should have been called since router didn't find the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_called()
|
||||
assert "api_key" not in _serialize_litellm_params(response_vs.get("litellm_params"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2425,9 +2190,7 @@ async def test_new_vector_store_auto_resolves_from_router():
|
|||
}
|
||||
return mock_created_vector_store
|
||||
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(
|
||||
side_effect=mock_create
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(side_effect=mock_create)
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.add_vector_store_to_registry = MagicMock()
|
||||
|
|
|
|||
|
|
@ -1,10 +1,19 @@
|
|||
import pytest
|
||||
import litellm
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.vector_stores.transformation import AzureAIVectorStoreConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.vector_stores import (
|
||||
asearch as vector_store_asearch,
|
||||
)
|
||||
from litellm.vector_stores import (
|
||||
search as vector_store_search,
|
||||
asearch as vector_store_asearch,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -30,10 +39,108 @@ async def test_basic_search_vector_store(sync_mode):
|
|||
if sync_mode:
|
||||
response = vector_store_search(query=default_query, **base_request_args)
|
||||
else:
|
||||
response = await vector_store_asearch(
|
||||
query=default_query, **base_request_args
|
||||
)
|
||||
response = await vector_store_asearch(query=default_query, **base_request_args)
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Skipping test due to litellm.InternalServerError")
|
||||
|
||||
print("litellm response=", json.dumps(response, indent=4, default=str))
|
||||
|
||||
|
||||
class RecordingEmbeddingExecutor:
|
||||
def __init__(self, response):
|
||||
self.response = response
|
||||
self.calls = []
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
|
||||
ALIAS_QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
ALIAS_EMBEDDING_RESPONSE = EmbeddingResponse(
|
||||
data=[{"embedding": ALIAS_QUERY_VECTOR, "index": 0, "object": "embedding"}]
|
||||
)
|
||||
STORE_EMBEDDINGS_URL = "https://embedding.example/v1/embeddings"
|
||||
|
||||
|
||||
def _transform_kwargs(executor):
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
return {
|
||||
"vector_store_id": "my-vector-index",
|
||||
"query": "what is azure search?",
|
||||
"vector_store_search_optional_params": {"top_k": 2},
|
||||
"api_base": "https://azure-kb-search.search.windows.net",
|
||||
"litellm_logging_obj": logging_obj,
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "multilingual-e5-large",
|
||||
"azure_search_vector_field": "embedding",
|
||||
},
|
||||
"embedding_executor": executor,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_uses_injected_executor_without_embedding_config(respx_mock: respx.MockRouter):
|
||||
executor = RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE)
|
||||
config = AzureAIVectorStoreConfig()
|
||||
transform_kwargs = _transform_kwargs(executor)
|
||||
|
||||
url, sync_body = config.transform_search_vector_store_request(**transform_kwargs)
|
||||
_, async_body = await config.atransform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
assert respx_mock.calls.call_count == 0
|
||||
assert executor.calls == [("multilingual-e5-large", "what is azure search?", {})] * 2
|
||||
assert (
|
||||
url == "https://azure-kb-search.search.windows.net/indexes/my-vector-index/docs/search?api-version=2024-07-01"
|
||||
)
|
||||
assert sync_body == async_body
|
||||
assert sync_body["vectorQueries"] == [
|
||||
{"vector": ALIAS_QUERY_VECTOR, "fields": "embedding", "kind": "vector", "k": 2}
|
||||
]
|
||||
assert sync_body["top"] == 2
|
||||
logging_details = transform_kwargs["litellm_logging_obj"].model_call_details
|
||||
assert logging_details["embedding_model"] == "multilingual-e5-large"
|
||||
assert logging_details["top_k"] == 2
|
||||
|
||||
|
||||
def test_transform_falls_back_to_sdk_embedding_without_executor(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = respx_mock.post(STORE_EMBEDDINGS_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": ALIAS_QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
transform_kwargs = _transform_kwargs(None)
|
||||
transform_kwargs["litellm_params"] = {
|
||||
"litellm_embedding_model": "openai/text-embedding-3-small",
|
||||
"litellm_embedding_config": {"api_base": "https://embedding.example/v1", "api_key": "store-key"},
|
||||
}
|
||||
|
||||
_, body = AzureAIVectorStoreConfig().transform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer store-key"
|
||||
assert json.loads(embedding_request.read())["input"] == ["what is azure search?"]
|
||||
assert body["vectorQueries"][0]["vector"] == ALIAS_QUERY_VECTOR
|
||||
assert body["vectorQueries"][0]["fields"] == "contentVector"
|
||||
|
||||
|
||||
def test_transform_requires_embedding_model():
|
||||
transform_kwargs = _transform_kwargs(RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE))
|
||||
transform_kwargs["litellm_params"] = {"litellm_embedding_config": {"api_key": "store-key"}}
|
||||
|
||||
with pytest.raises(ValueError, match="litellm_embedding_model is required"):
|
||||
AzureAIVectorStoreConfig().transform_search_vector_store_request(**transform_kwargs)
|
||||
|
|
|
|||
|
|
@ -3,16 +3,19 @@ Tests for Milvus Vector Store
|
|||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.llms.milvus.vector_stores.transformation import MilvusVectorStoreConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.vector_stores import asearch as vector_store_asearch
|
||||
from litellm.vector_stores import search as vector_store_search
|
||||
|
||||
|
||||
# Mock response from actual Milvus API
|
||||
MOCK_MILVUS_SEARCH_RESPONSE = {
|
||||
"code": 0,
|
||||
|
|
@ -98,7 +101,7 @@ class TestMilvusVectorStore:
|
|||
mock_response.json.return_value = MOCK_MILVUS_SEARCH_RESPONSE
|
||||
mock_response.text = json.dumps(MOCK_MILVUS_SEARCH_RESPONSE)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
with patch("litellm.aembedding", new_callable=AsyncMock) as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
|
|
@ -147,16 +150,10 @@ class TestMilvusVectorStore:
|
|||
else:
|
||||
# Fallback: check for json kwarg or in args
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
|
||||
assert (
|
||||
request_data is not None
|
||||
), f"Could not extract request data. Call args: {call_args}"
|
||||
assert request_data is not None, f"Could not extract request data. Call args: {call_args}"
|
||||
print("Request data:", json.dumps(request_data, indent=2, default=str))
|
||||
|
||||
# Validate request structure
|
||||
|
|
@ -213,9 +210,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Make the search request
|
||||
|
|
@ -252,16 +247,10 @@ class TestMilvusVectorStore:
|
|||
else:
|
||||
# Fallback: check for json kwarg or in args
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
|
||||
assert (
|
||||
request_data is not None
|
||||
), f"Could not extract request data. Call args: {call_args}"
|
||||
assert request_data is not None, f"Could not extract request data. Call args: {call_args}"
|
||||
|
||||
# Validate request structure
|
||||
assert "collectionName" in request_data
|
||||
|
|
@ -316,11 +305,7 @@ class TestMilvusVectorStore:
|
|||
if request_data_str:
|
||||
return json.loads(request_data_str)
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
return request_data
|
||||
|
||||
|
|
@ -334,9 +319,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -375,9 +358,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -413,9 +394,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -492,3 +471,247 @@ if __name__ == "__main__":
|
|||
test.test_basic_search_with_mock_sync()
|
||||
|
||||
print("\n✅ All mock tests passed!")
|
||||
|
||||
|
||||
class RecordingEmbeddingExecutor:
|
||||
def __init__(self, response):
|
||||
self.response = response
|
||||
self.calls = []
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
|
||||
ALIAS_QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
ALIAS_EMBEDDING_RESPONSE = EmbeddingResponse(
|
||||
data=[{"embedding": ALIAS_QUERY_VECTOR, "index": 0, "object": "embedding"}]
|
||||
)
|
||||
OPENAI_EMBEDDINGS_URL = "https://api.openai.com/v1/embeddings"
|
||||
MILVUS_SEARCH_URL = "https://milvus.example/v2/vectordb/entities/search"
|
||||
ALIAS_SEARCH_KWARGS = {
|
||||
"query": "what is machine learning?",
|
||||
"vector_store_id": "book_2",
|
||||
"custom_llm_provider": "milvus",
|
||||
"api_base": "https://milvus.example",
|
||||
"api_key": "mock_milvus_api_key",
|
||||
"litellm_embedding_model": "multilingual-e5-large",
|
||||
"milvus_text_field": "book_intro_text",
|
||||
}
|
||||
|
||||
|
||||
def _alias_router():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "multilingual-e5-large",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _mock_embedding_route(respx_mock: respx.MockRouter) -> respx.Route:
|
||||
return respx_mock.post(OPENAI_EMBEDDINGS_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": ALIAS_QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _mock_search_route(respx_mock: respx.MockRouter) -> respx.Route:
|
||||
return respx_mock.post(MILVUS_SEARCH_URL).mock(return_value=httpx.Response(200, json=MOCK_MILVUS_SEARCH_RESPONSE))
|
||||
|
||||
|
||||
def _assert_alias_resolved(embedding_route: respx.Route, search_route: respx.Route, response):
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer deployment-key"
|
||||
embedding_body = json.loads(embedding_request.read())
|
||||
assert embedding_body["model"] == "text-embedding-3-small"
|
||||
assert embedding_body["input"] == ["what is machine learning?"]
|
||||
search_request = search_route.calls.last.request
|
||||
assert search_request.headers["authorization"] == "Bearer mock_milvus_api_key"
|
||||
assert json.loads(search_request.read())["data"] == [ALIAS_QUERY_VECTOR]
|
||||
assert len(response["data"]) == len(MOCK_MILVUS_SEARCH_RESPONSE["data"])
|
||||
assert response["data"][0]["content"][0]["text"] == MOCK_MILVUS_SEARCH_RESPONSE["data"][0]["book_intro_text"]
|
||||
|
||||
|
||||
def test_router_search_resolves_bare_embedding_alias_sync(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = _alias_router().vector_store_search(**ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_search_resolves_bare_embedding_alias_async(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await _alias_router().avector_store_search(**ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
def test_sdk_search_with_router_kwarg_resolves_bare_embedding_alias_sync(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = litellm.vector_stores.search(router=_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_resolves_bare_embedding_alias_async(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await litellm.vector_stores.asearch(router=_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
def _team_alias_router():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "team-a-embedder",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
"model_info": {"team_id": "team-a", "team_public_model_name": "multilingual-e5-large"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_resolves_team_alias_from_request_metadata(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await litellm.vector_stores.asearch(
|
||||
router=_team_alias_router(), metadata={"user_api_key_team_id": "team-a"}, **ALIAS_SEARCH_KWARGS
|
||||
)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_rejects_team_alias_without_team_metadata(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
_mock_search_route(respx_mock)
|
||||
|
||||
with pytest.raises(litellm.APIConnectionError):
|
||||
await litellm.vector_stores.asearch(router=_team_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
assert embedding_route.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_uses_injected_executor_without_embedding_config(respx_mock: respx.MockRouter):
|
||||
executor = RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE)
|
||||
config = MilvusVectorStoreConfig()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
transform_kwargs = {
|
||||
"vector_store_id": "book_2",
|
||||
"query": ["what is", "milvus?"],
|
||||
"vector_store_search_optional_params": {"limit": 3},
|
||||
"api_base": "https://milvus.example",
|
||||
"litellm_logging_obj": logging_obj,
|
||||
"litellm_params": {"litellm_embedding_model": "multilingual-e5-large", "milvus_db_name": "docs"},
|
||||
"embedding_executor": executor,
|
||||
}
|
||||
|
||||
url, sync_body = config.transform_search_vector_store_request(**transform_kwargs)
|
||||
_, async_body = await config.atransform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
assert respx_mock.calls.call_count == 0
|
||||
assert executor.calls == [("multilingual-e5-large", "what is milvus?", {})] * 2
|
||||
assert url == MILVUS_SEARCH_URL
|
||||
assert sync_body == async_body
|
||||
assert sync_body == {
|
||||
"collectionName": "book_2",
|
||||
"data": [ALIAS_QUERY_VECTOR],
|
||||
"annsField": "book_intro_vector",
|
||||
"limit": 3,
|
||||
"dbName": "docs",
|
||||
}
|
||||
assert logging_obj.model_call_details["input"] == "what is milvus?"
|
||||
assert logging_obj.model_call_details["embedding_model"] == "multilingual-e5-large"
|
||||
|
||||
|
||||
def test_transform_falls_back_to_sdk_embedding_without_executor_or_config(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "env-key")
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
|
||||
_, body = MilvusVectorStoreConfig().transform_search_vector_store_request(
|
||||
vector_store_id="book_2",
|
||||
query="q",
|
||||
vector_store_search_optional_params={},
|
||||
api_base="https://milvus.example",
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small"},
|
||||
)
|
||||
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer env-key"
|
||||
assert json.loads(embedding_request.read())["input"] == ["q"]
|
||||
assert body["data"] == [ALIAS_QUERY_VECTOR]
|
||||
|
||||
|
||||
def test_transform_requires_embedding_model():
|
||||
with pytest.raises(ValueError, match="litellm_embedding_model is required"):
|
||||
MilvusVectorStoreConfig().transform_search_vector_store_request(
|
||||
vector_store_id="book_2",
|
||||
query="q",
|
||||
vector_store_search_optional_params={},
|
||||
api_base="https://milvus.example",
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_params={"litellm_embedding_config": {"api_key": "store-key"}},
|
||||
embedding_executor=RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE),
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue