mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(vector-store): route embeddings through router
This commit is contained in:
parent
bfa5eac76b
commit
5635811726
11 changed files with 469 additions and 769 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,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
22
uv.lock
generated
|
|
@ -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]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue