fix(vector-store): route embeddings through router

This commit is contained in:
Yujong Lee 2026-09-01 12:15:03 -07:00
parent bfa5eac76b
commit 5635811726
11 changed files with 469 additions and 769 deletions

View file

@ -1,10 +1,14 @@
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 typing import TYPE_CHECKING, Any, NoReturn, Protocol, runtime_checkable
import httpx
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import (
VECTOR_STORE_OPENAI_PARAMS,
BaseVectorStoreAuthCredentials,
@ -17,6 +21,7 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
from ..chat.transformation import BaseLLMException as _BaseLLMException
@ -27,6 +32,58 @@ 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
)
@dataclass(frozen=True, slots=True)
class RouterVectorStoreEmbeddingExecutor:
router: Router
metadata: Mapping[str, object]
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
if configuration:
return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, configuration)
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
metadata=dict(self.metadata), # mutable-ok: Router metadata requires a concrete dict
)
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
if configuration:
return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, configuration)
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
metadata=dict(self.metadata), # mutable-ok: Router metadata requires a concrete dict
)
class BaseVectorStoreConfig:
def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]:
return []
@ -172,6 +229,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
@ -184,6 +242,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

View file

@ -70,6 +70,7 @@ from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeech
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
BaseVectorStoreConfig,
VectorStoreEmbeddingExecutor,
)
from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
@ -9683,6 +9684,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,
@ -9702,6 +9704,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,
)
@ -9797,6 +9800,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,
@ -9812,6 +9816,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,
@ -9831,6 +9836,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,
)

View file

@ -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

View file

@ -14,9 +14,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,
@ -65,32 +62,9 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
data["litellm_credential_name"] = vector_store_to_run.get("litellm_credential_name")
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
# Resolve ``litellm_embedding_config`` here, at request-handling
# time, instead of at row-creation time. The resolved
# ``api_key`` / ``api_base`` / ``api_version`` lives only in
# this per-request ``data`` dict and is never persisted.
# Legacy rows that carry a resolved config are refreshed when the
# embedding model is an alias so the provider-qualified model is used.
embedding_model: Final = litellm_params.get("litellm_embedding_model")
if embedding_model:
from litellm.proxy.proxy_server import prisma_client
embedding_resolution: Final = await _resolve_embedding_config(
embedding_model=embedding_model, prisma_client=prisma_client
)
if embedding_resolution:
resolved_model, resolved_config = embedding_resolution
# Build a fresh dict via spread instead of mutating
# ``litellm_params`` in place — the registry hands back
# a reference to its cached object, so an in-place
# update would persist the resolved cleartext into the
# in-memory cache for the lifetime of the process.
litellm_params = {
**litellm_params,
"litellm_embedding_model": resolved_model,
"litellm_embedding_config": resolved_config,
}
litellm_params: Final = (
vector_store_to_run.get("litellm_params", {}) or {}
) # mutable-ok: request execution merges persisted params into a mutable body
data.update(litellm_params)
return data

View file

@ -10,7 +10,7 @@ All /vector_store management endpoints
import copy
import json
from typing import TYPE_CHECKING, Any, Final, TypeAlias
from typing import TYPE_CHECKING, Any, Final
from fastapi import APIRouter, Depends, HTTPException
@ -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,
@ -49,7 +43,6 @@ from litellm.types.vector_stores import (
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
router: Final = APIRouter()
EmbeddingResolution: TypeAlias = tuple[str, dict[str, object]]
def _vector_store_table(prisma_client: "PrismaClient") -> "TableActions[_VectorStoreRow]":
@ -65,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:
"""
@ -156,264 +127,6 @@ async def _fetch_and_authorize_vector_store(
return typed
def _provider_qualified_embedding_model(
fallback: str,
model: object,
custom_llm_provider: object,
) -> str:
if not isinstance(model, str) or not model:
return fallback
if "/" in model or not isinstance(custom_llm_provider, str) or not custom_llm_provider:
return model
return f"{custom_llm_provider}/{model}"
def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> EmbeddingResolution | 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:
Provider-qualified model and its connection config if found, otherwise None
"""
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
resolved_model: Final = _provider_qualified_embedding_model(
fallback=embedding_model,
model=getattr(litellm_params, "model", None),
custom_llm_provider=getattr(litellm_params, "custom_llm_provider", None),
)
# 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 (
resolved_model,
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"
) -> EmbeddingResolution | 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:
Provider-qualified model and its connection config if found, otherwise None
"""
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()),
)
resolved_model: Final = _provider_qualified_embedding_model(
fallback=embedding_model,
model=decrypted_params.get("model"),
custom_llm_provider=decrypted_params.get("custom_llm_provider"),
)
return (
resolved_model,
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
) -> EmbeddingResolution | 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:
Provider-qualified model and its connection config if found, otherwise None
"""
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
########################################################

View file

@ -84,6 +84,9 @@ 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,
)
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
@ -6319,6 +6322,34 @@ class Router:
client: object | None = None,
**kwargs,
):
if call_type == "vector_store_search":
metadata: Final = self._vector_store_request_metadata(kwargs)
provider_kwargs: Final = (
{
"custom_llm_provider": custom_llm_provider
} # mutable-ok: provider kwargs are expanded into the request
if custom_llm_provider is not None
else MappingProxyType({})
)
search_kwargs: Final = { # mutable-ok: the routed request requires dynamic keyword arguments
**kwargs,
**provider_kwargs,
"_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor(
router=self,
metadata=metadata,
),
}
model: Final = search_kwargs.get("model")
if isinstance(model, str) and model:
routed_kwargs: Final = { # mutable-ok: model must be removed before expanding routed kwargs
key: value for key, value in search_kwargs.items() if key != "model"
}
return self._generic_api_call_with_fallbacks(
model=model,
original_function=original_function,
**routed_kwargs,
)
return original_function(**search_kwargs)
return self._generic_api_call_with_fallbacks(original_function=original_function, **kwargs)
return sync_wrapper
@ -6512,10 +6543,21 @@ 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,
**kwargs,
**vector_store_kwargs,
)
elif call_type in ("afile_delete", "afile_content"):
return await self._ageneric_api_call_with_fallbacks(
@ -6551,6 +6593,18 @@ class Router:
return async_wrapper
@staticmethod
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 cast( # cast-ok: isinstance validates the runtime dict boundary
"dict[str, object]", litellm_metadata
)
metadata: Final = kwargs.get("metadata")
if isinstance(metadata, dict):
return cast("dict[str, object]", metadata) # cast-ok: isinstance validates the runtime dict boundary
return MappingProxyType({})
async def _init_vector_store_api_endpoints(
self,
original_function: Callable,

View file

@ -15,6 +15,10 @@ 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 (
LiteLLMVectorStoreEmbeddingExecutor,
VectorStoreEmbeddingExecutor,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
@ -35,6 +39,14 @@ base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
def _direct_vector_store_embedding_executor(value: object) -> VectorStoreEmbeddingExecutor:
if value is None:
return LiteLLMVectorStoreEmbeddingExecutor()
if isinstance(value, VectorStoreEmbeddingExecutor):
return value
raise TypeError("Invalid direct vector store embedding executor")
def mock_vector_store_search_response(
mock_results: list[VectorStoreSearchResult] | None = None,
):
@ -285,7 +297,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)
)
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()
@ -308,6 +325,7 @@ async def asearch(
extra_body=extra_body,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
_direct_vector_store_embedding_executor=embedding_executor,
**kwargs,
)
@ -363,12 +381,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)
)
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:
@ -445,6 +467,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,

View file

@ -5,17 +5,107 @@ 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
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm import Router
from litellm.llms.base_llm.vector_store.transformation import (
LiteLLMVectorStoreEmbeddingExecutor,
RouterVectorStoreEmbeddingExecutor,
)
from litellm.types.utils import EmbeddingResponse
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):
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
sdk_executor = LiteLLMVectorStoreEmbeddingExecutor()
with (
patch("litellm.embedding", return_value=response) as embedding,
patch("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
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("litellm.embedding", return_value=response) as explicit_embedding:
assert router_executor.embed("openai/model", "query", {"api_key": "store-key"}) is response
explicit_embedding.assert_called_once_with(model="openai/model", input=["query"], api_key="store-key")
mock_router.embedding.assert_called_once()
with patch("litellm.aembedding", new=AsyncMock(return_value=response)) as explicit_aembedding:
assert await router_executor.aembed("openai/model", "query", {"api_key": "store-key"}) is response
explicit_aembedding.assert_awaited_once_with(model="openai/model", input=["query"], api_key="store-key")
def test_embedding_with_deployment_specific_headers(self):
"""
Test that deployment-specific headers are propagated.

View file

@ -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()

View file

@ -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,98 @@ 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())
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()
wrapped_create = router.factory_function(create_original, call_type="vector_store_create")
with patch.object( # test-quality-ok: fallback dispatch is the boundary this wrapper delegates to
router, "_generic_api_call_with_fallbacks", return_value="created"
) as fallback:
assert wrapped_create(name="store") == "created"
fallback.assert_called_once_with(original_function=create_original, 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
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_called_once_with(model="openai/model", input=["query"], api_key="store-key")
explicit_aembedding.assert_awaited_once_with(model="openai/model", input=["query"], api_key="store-key")
@pytest.mark.asyncio
@ -82,10 +162,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 +177,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,89 +615,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": "multilingual-e5-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=("azure/multilingual-e5-large", 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/multilingual-e5-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_preserves_legacy_embedding_config_when_model_not_resolved():
"""A vector store row created by an older proxy version may already
carry a fully-resolved ``litellm_embedding_config`` in its persisted
``litellm_params``. Preserve it when the model cannot be resolved."""
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(return_value=None)
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_awaited_once()
class TestCheckVectorStorePermission:
"""Test suite for check_vector_store_permission function."""
@ -2001,60 +2055,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 = {
"model": "openai/text-embedding-3-small",
"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
resolved_model, resolved_config = result
assert resolved_model == "openai/text-embedding-3-small"
assert resolved_config["api_key"] == "test-api-key"
assert resolved_config["api_base"] == "https://api.openai.com"
assert resolved_config["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
@ -2071,14 +2072,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
@ -2089,10 +2082,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 = {}
@ -2113,280 +2102,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_litellm_params.model = "text-embedding-3-small"
mock_litellm_params.custom_llm_provider = "openai"
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
resolved_model, resolved_config = result
assert resolved_model == "openai/text-embedding-3-small"
assert resolved_config["api_key"] == "config-api-key"
assert resolved_config["api_base"] == "https://config-api-base.com"
assert resolved_config["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_litellm_params.model = "text-embedding-3-large"
mock_litellm_params.custom_llm_provider = "azure"
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
resolved_model, resolved_config = result
assert resolved_model == "azure/text-embedding-3-large"
assert resolved_config["api_key"] == "azure-api-key"
assert resolved_config["api_base"] == "https://azure-endpoint.openai.azure.com"
assert resolved_config["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_litellm_params.model = "text-embedding-3-small"
mock_litellm_params.custom_llm_provider = "openai"
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
resolved_model, resolved_config = result
assert resolved_model == "openai/text-embedding-3-small"
assert resolved_config["api_key"] == "resolved-from-env"
assert resolved_config["api_base"] == "https://direct-url.com"
assert "api_version" not in resolved_config
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_litellm_params.model = "text-embedding-3-small"
mock_litellm_params.custom_llm_provider = "openai"
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
resolved_model, resolved_config = result
assert resolved_model == "openai/text-embedding-3-small"
assert resolved_config["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 = {
"model": "openai/text-embedding-3-small",
"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
resolved_model, resolved_config = result
assert resolved_model == "openai/text-embedding-3-small"
assert resolved_config["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
@ -2445,9 +2175,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()

22
uv.lock generated
View file

@ -9441,19 +9441,19 @@ wheels = [
[[package]]
name = "tornado"
version = "6.5.7"
version = "6.5.8"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/64/24/95ec527ad67b76d59299e5465b3935d05e4294b7e0290a3924b7487df30b/tornado-6.5.7.tar.gz", hash = "sha256:66c513a76cda70d53907bc27cf1447557699c2e95aa48ba27a442ff61c3ddfc2", size = 519252, upload-time = "2026-06-08T17:34:51.232Z" }
sdist = { url = "https://files.pythonhosted.org/packages/10/d3/343e5bb989d6515b1646cf3d40135d73f3d5e45339bded401b56cdac24dd/tornado-6.5.8.tar.gz", hash = "sha256:9452e1b208a8bd771e2cb1f2ff564985b9b214bdebbe622793e1799e0a6bd23f", size = 520493, upload-time = "2026-08-07T02:12:42.971Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/02/dc/c7043cab6fed8ae159fc1923ce829ada35c4dbd797d408a43858ffaf9639/tornado-6.5.7-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:148b2eb15c2c765a50796172c1e499649b35f30d2e3c3d3e15913cfa56bfb163", size = 448543, upload-time = "2026-06-08T17:34:38.052Z" },
{ url = "https://files.pythonhosted.org/packages/92/4f/090b1431e5a43df696feceffc268c5383cc079ecb5f08ce58f917109aafe/tornado-6.5.7-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9da38de27f1da3b78a966f0dae12b5a1ea9afe72ca805d84ff06508272ddf100", size = 446707, upload-time = "2026-06-08T17:34:39.594Z" },
{ url = "https://files.pythonhosted.org/packages/37/d8/ef374952fd5da67d4463122c2b8e5a96536ec10b4b339254c6dcde81d01c/tornado-6.5.7-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:8d759e71906ee783f8867b93bf26a265743da4c1e2f4a018464c1ba019862972", size = 449774, upload-time = "2026-06-08T17:34:41.204Z" },
{ url = "https://files.pythonhosted.org/packages/35/37/d434c73f4c6e014b745b9b37085f34f40c022f007efff3d7fe65991899f3/tornado-6.5.7-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8a46347a18f23fb92b396beebe0fb78f61dda0cc302445202c16203d8a18848b", size = 450745, upload-time = "2026-06-08T17:34:42.531Z" },
{ url = "https://files.pythonhosted.org/packages/b6/2b/56b9aff361d7f1ab728a805ec7d7ea835f8807afa9f5cc690ea0e630efb9/tornado-6.5.7-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:7778b30bef919231265e91c69963ce0f49a1e9c07ac900bbe75b19ce2575ba92", size = 450578, upload-time = "2026-06-08T17:34:43.787Z" },
{ url = "https://files.pythonhosted.org/packages/02/30/a7444fb23aa76860a14198fab96ac79f1866b0a6e19e26c4381b0938e50f/tornado-6.5.7-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:e726f0c75da7726eec023aa62751ff8878bd2737e34fbdd33b1ae5897d2200f5", size = 449985, upload-time = "2026-06-08T17:34:45.326Z" },
{ url = "https://files.pythonhosted.org/packages/5c/42/5f0e56c01e8d9d36f4e23f367b85ae6cae0c1ecddd5e6977d8388ad27488/tornado-6.5.7-cp39-abi3-win32.whl", hash = "sha256:f8de3bf12d3efdd0cbe7c8887868198f8a91415e3f29fcf258d9b8eb7b1d9ae4", size = 451047, upload-time = "2026-06-08T17:34:46.784Z" },
{ url = "https://files.pythonhosted.org/packages/c9/a4/b393076ffb21b469eec5b328a0534cf03a3b90bfc6b1f09507cdd075d938/tornado-6.5.7-cp39-abi3-win_amd64.whl", hash = "sha256:de942f843533a039ef9fa3d9c88c7cd8a7c94553fb5ad0154270989b3d99a2c4", size = 451485, upload-time = "2026-06-08T17:34:48.248Z" },
{ url = "https://files.pythonhosted.org/packages/71/2e/7b1c769803121b809112cf9a00681c472eae1d80e32d7ec0e0bd61d0d0e1/tornado-6.5.7-cp39-abi3-win_arm64.whl", hash = "sha256:ff934fce95643af5f11efdae618eaa73d469dc588641e5c8d19295a0c65c4796", size = 450506, upload-time = "2026-06-08T17:34:49.702Z" },
{ url = "https://files.pythonhosted.org/packages/f2/d5/007086fd8df5489338e204f65adce33fd4f21a4999dbb2b9cff2f897b5f4/tornado-6.5.8-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:cc6aa787d7cfab7c3d35189dc7a56fbd2399a569624c730c6b55b3d6531d0403", size = 449487, upload-time = "2026-08-07T02:12:28.682Z" },
{ url = "https://files.pythonhosted.org/packages/70/c8/5a24a99495903f594f6a199dd7beead1cbc0a13e2cb9102727bcaaf2a997/tornado-6.5.8-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9715b5eb79735b2bcd454ce216a9275b7c0470e64ea1bf5742f78b2f72b26eeb", size = 447649, upload-time = "2026-08-07T02:12:30.306Z" },
{ url = "https://files.pythonhosted.org/packages/6e/de/f2e733f386b85962d1b1dc82cd63d169b5b4580062b35397eac9244a41fe/tornado-6.5.8-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:547d63f450d570c14fe0e8db2cfb14c9bbd1c2503b4a6612586267955aa47b58", size = 450707, upload-time = "2026-08-07T02:12:31.95Z" },
{ url = "https://files.pythonhosted.org/packages/0b/94/20efeee9a01c141e9ac47c397f81679dfda24b32768fc4fff24e76d36c2c/tornado-6.5.8-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7e2360a0ffbe145eca8af0b19cb7203d79b1a98dd4cccdd6b368f6f49c2e3808", size = 451677, upload-time = "2026-08-07T02:12:33.512Z" },
{ url = "https://files.pythonhosted.org/packages/42/ec/a96ccb8ccf0de2b7bc2c5fa1608a4803735018242e90c4882365a9fd418f/tornado-6.5.8-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:5d242290bdf7ab3151bc1065fdd75c0dcc21cbc7b49f22a4c56329c2d6566d22", size = 451510, upload-time = "2026-08-07T02:12:35.346Z" },
{ url = "https://files.pythonhosted.org/packages/29/b5/93185859245ad3f00e62175f29607346788b696369347f0146e0421286bb/tornado-6.5.8-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7b94ff0e128fe0542f3bd331fb44d06260fc4ac16881545159f34ef08aad4195", size = 450917, upload-time = "2026-08-07T02:12:36.963Z" },
{ url = "https://files.pythonhosted.org/packages/97/cf/fe33cf062834487d34d1559746a4a12521033c22645b6d74d4bca702e018/tornado-6.5.8-cp39-abi3-win32.whl", hash = "sha256:67832909c4779c64942380cb5f044a5c6163d00831472d80e25e115de9917836", size = 451952, upload-time = "2026-08-07T02:12:38.512Z" },
{ url = "https://files.pythonhosted.org/packages/cb/e1/468ad54333e92ccb62627e62cb88e5fc14a2171daa67ed47b1b8542d5b86/tornado-6.5.8-cp39-abi3-win_amd64.whl", hash = "sha256:11881db6b7c168494be2c2d12e65931451bdf7ee718535418ae1d8855dd5a0ee", size = 452391, upload-time = "2026-08-07T02:12:39.971Z" },
{ url = "https://files.pythonhosted.org/packages/ad/3e/cd5e4f06e34cde33b8ef66cf36aa2b5ad46354cc1af7d2136bbe365fee1d/tornado-6.5.8-cp39-abi3-win_arm64.whl", hash = "sha256:68a7468c7e289f8514d7d664101753903217eff1bb6822c6b5994a0b5f5bcb26", size = 451411, upload-time = "2026-08-07T02:12:41.469Z" },
]
[[package]]